mirror of https://github.com/apache/cassandra
Implement proper sandboxing for UDFs
patch by Robert Stupp; reviewed by T Jake Luciani for CASSANDRA-9402
This commit is contained in:
parent
48fc0b8d0d
commit
5790b4a44b
|
|
@ -1,4 +1,5 @@
|
|||
3.0
|
||||
* Implement proper sandboxing for UDFs (CASSANDRA-9402)
|
||||
* Allow extra schema definitions in cassandra-stress yaml (CASSANDRA-9850)
|
||||
* Metrics should use up to date nomenclature (CASSANDRA-9448)
|
||||
* Change CREATE/ALTER TABLE syntax for compression (CASSANDRA-8384)
|
||||
|
|
|
|||
3
NEWS.txt
3
NEWS.txt
|
|
@ -18,7 +18,8 @@ using the provided 'sstableupgrade' tool.
|
|||
|
||||
Upgrading
|
||||
---------
|
||||
_ New SSTable version 'la' with improved bloom-filter false-positive handling
|
||||
- User defined functions are now by default enabled.
|
||||
- New SSTable version 'la' with improved bloom-filter false-positive handling
|
||||
compared to previous version 'ka' used in 2.2 and 2.1. Running sstableupgrade
|
||||
is not necessary but recommended.
|
||||
- Before upgrading to 3.0, make sure that your cluster is in complete agreement
|
||||
|
|
|
|||
|
|
@ -872,12 +872,9 @@ inter_dc_tcp_nodelay: false
|
|||
tracetype_query_ttl: 86400
|
||||
tracetype_repair_ttl: 604800
|
||||
|
||||
# UDFs (user defined functions) are disabled by default.
|
||||
# As of Cassandra 2.2, there is no security manager or anything else in place that
|
||||
# prevents execution of evil code. CASSANDRA-9402 will fix this issue for Cassandra 3.0.
|
||||
# This will inherently be backwards-incompatible with any 2.2 UDF that perform insecure
|
||||
# operations such as opening a socket or writing to the filesystem.
|
||||
enable_user_defined_functions: false
|
||||
# UDFs (user defined functions) are enabled by default.
|
||||
# As of Cassandra 3.0 there is a sandbox in place that prevents execution of evil code.
|
||||
enable_user_defined_functions: true
|
||||
|
||||
# The default Windows kernel timer and scheduling resolution is 15.6ms for power conservation.
|
||||
# Lowering this value on Windows can provide much tighter latency and better throughput, however
|
||||
|
|
|
|||
|
|
@ -203,7 +203,8 @@ public abstract class AbstractTracingAwareExecutorService implements TracingAwar
|
|||
|
||||
public T get(long timeout, TimeUnit unit) throws InterruptedException, ExecutionException, TimeoutException
|
||||
{
|
||||
await(timeout, unit);
|
||||
if (!await(timeout, unit))
|
||||
throw new TimeoutException();
|
||||
Object result = this.result;
|
||||
if (failure)
|
||||
throw new ExecutionException((Throwable) result);
|
||||
|
|
|
|||
|
|
@ -30,6 +30,8 @@ public class NamedThreadFactory implements ThreadFactory
|
|||
{
|
||||
protected final String id;
|
||||
private final int priority;
|
||||
private final ClassLoader contextClassLoader;
|
||||
private final ThreadGroup threadGroup;
|
||||
protected final AtomicInteger n = new AtomicInteger(1);
|
||||
|
||||
public NamedThreadFactory(String id)
|
||||
|
|
@ -39,17 +41,25 @@ public class NamedThreadFactory implements ThreadFactory
|
|||
|
||||
public NamedThreadFactory(String id, int priority)
|
||||
{
|
||||
this(id, priority, null, null);
|
||||
}
|
||||
|
||||
public NamedThreadFactory(String id, int priority, ClassLoader contextClassLoader, ThreadGroup threadGroup)
|
||||
{
|
||||
this.id = id;
|
||||
this.priority = priority;
|
||||
this.contextClassLoader = contextClassLoader;
|
||||
this.threadGroup = threadGroup;
|
||||
}
|
||||
|
||||
public Thread newThread(Runnable runnable)
|
||||
{
|
||||
String name = id + ":" + n.getAndIncrement();
|
||||
Thread thread = new Thread(runnable, name);
|
||||
Thread thread = new Thread(threadGroup, runnable, name);
|
||||
thread.setPriority(priority);
|
||||
thread.setDaemon(true);
|
||||
if (contextClassLoader != null)
|
||||
thread.setContextClassLoader(contextClassLoader);
|
||||
return thread;
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -256,7 +256,37 @@ public class Config
|
|||
|
||||
public int windows_timer_interval = 0;
|
||||
|
||||
public boolean enable_user_defined_functions = false;
|
||||
public boolean enable_user_defined_functions = true;
|
||||
/**
|
||||
* Optionally disable asynchronous UDF execution.
|
||||
* Disabling asynchronous UDF execution also implicitly disables the security-manager!
|
||||
* By default, async UDF execution is enabled to be able to detect UDFs that run too long / forever and be
|
||||
* able to fail fast - i.e. stop the Cassandra daemon, which is currently the only appropriate approach to
|
||||
* "tell" a user that there's something really wrong with the UDF.
|
||||
* When you disable async UDF execution, users MUST pay attention to read-timeouts since these may indicate
|
||||
* UDFs that run too long or forever - and this can destabilize the cluster.
|
||||
*/
|
||||
public boolean enable_user_defined_functions_threads = true;
|
||||
/**
|
||||
* Time in milliseconds after a warning will be emitted to the log and to the client that a UDF runs too long.
|
||||
* (Only valid, if enable_user_defined_functions_threads==true)
|
||||
*/
|
||||
public long user_defined_function_warn_timeout = 500;
|
||||
/**
|
||||
* Time in milliseconds after a fatal UDF run-time situation is detected and action according to
|
||||
* user_function_timeout_policy will take place.
|
||||
* (Only valid, if enable_user_defined_functions_threads==true)
|
||||
*/
|
||||
public long user_defined_function_fail_timeout = 1500;
|
||||
/**
|
||||
* Defines what to do when a UDF ran longer than user_defined_function_fail_timeout.
|
||||
* Possible options are:
|
||||
* - 'die' - i.e. it is able to emit a warning to the client before the Cassandra Daemon will shut down.
|
||||
* - 'die_immediate' - shut down C* daemon immediately (effectively prevent the chance that the client will receive a warning).
|
||||
* - 'ignore' - just log - the most dangerous option.
|
||||
* (Only valid, if enable_user_defined_functions_threads==true)
|
||||
*/
|
||||
public UserFunctionTimeoutPolicy user_function_timeout_policy = UserFunctionTimeoutPolicy.die;
|
||||
|
||||
public static boolean getOutboundBindAny()
|
||||
{
|
||||
|
|
@ -321,6 +351,13 @@ public class Config
|
|||
die,
|
||||
}
|
||||
|
||||
public static enum UserFunctionTimeoutPolicy
|
||||
{
|
||||
ignore,
|
||||
die,
|
||||
die_immediate
|
||||
}
|
||||
|
||||
public static enum RequestSchedulerId
|
||||
{
|
||||
keyspace
|
||||
|
|
|
|||
|
|
@ -606,6 +606,14 @@ public class DatabaseDescriptor
|
|||
}
|
||||
if (seedProvider.getSeeds().size() == 0)
|
||||
throw new ConfigurationException("The seed provider lists no seeds.", false);
|
||||
|
||||
if (conf.user_defined_function_fail_timeout < 0)
|
||||
throw new ConfigurationException("user_defined_function_fail_timeout must not be negative", false);
|
||||
if (conf.user_defined_function_warn_timeout < 0)
|
||||
throw new ConfigurationException("user_defined_function_warn_timeout must not be negative", false);
|
||||
|
||||
if (conf.user_defined_function_fail_timeout < conf.user_defined_function_warn_timeout)
|
||||
throw new ConfigurationException("user_defined_function_warn_timeout must less than user_defined_function_fail_timeout", false);
|
||||
}
|
||||
|
||||
private static IEndpointSnitch createEndpointSnitch(String snitchClassName) throws ConfigurationException
|
||||
|
|
@ -1691,13 +1699,48 @@ public class DatabaseDescriptor
|
|||
return conf.otc_coalescing_window_us;
|
||||
}
|
||||
|
||||
public static int getWindowsTimerInterval()
|
||||
{
|
||||
return conf.windows_timer_interval;
|
||||
}
|
||||
|
||||
public static boolean enableUserDefinedFunctions()
|
||||
{
|
||||
return conf.enable_user_defined_functions;
|
||||
}
|
||||
|
||||
public static int getWindowsTimerInterval()
|
||||
public static boolean enableUserDefinedFunctionsThreads()
|
||||
{
|
||||
return conf.windows_timer_interval;
|
||||
return conf.enable_user_defined_functions_threads;
|
||||
}
|
||||
|
||||
public static long getUserDefinedFunctionWarnTimeout()
|
||||
{
|
||||
return conf.user_defined_function_warn_timeout;
|
||||
}
|
||||
|
||||
public static void setUserDefinedFunctionWarnTimeout(long userDefinedFunctionWarnTimeout)
|
||||
{
|
||||
conf.user_defined_function_warn_timeout = userDefinedFunctionWarnTimeout;
|
||||
}
|
||||
|
||||
public static long getUserDefinedFunctionFailTimeout()
|
||||
{
|
||||
return conf.user_defined_function_fail_timeout;
|
||||
}
|
||||
|
||||
public static void setUserDefinedFunctionFailTimeout(long userDefinedFunctionFailTimeout)
|
||||
{
|
||||
conf.user_defined_function_fail_timeout = userDefinedFunctionFailTimeout;
|
||||
}
|
||||
|
||||
public static Config.UserFunctionTimeoutPolicy getUserFunctionTimeoutPolicy()
|
||||
{
|
||||
return conf.user_function_timeout_policy;
|
||||
}
|
||||
|
||||
public static void setUserFunctionTimeoutPolicy(Config.UserFunctionTimeoutPolicy userFunctionTimeoutPolicy)
|
||||
{
|
||||
conf.user_function_timeout_policy = userFunctionTimeoutPolicy;
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@
|
|||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
package org.apache.cassandra.cql3.functions;
|
||||
|
||||
import java.io.ByteArrayOutputStream;
|
||||
|
|
@ -24,6 +25,16 @@ import java.lang.invoke.MethodHandle;
|
|||
import java.lang.invoke.MethodHandles;
|
||||
import java.lang.invoke.MethodType;
|
||||
import java.lang.reflect.InvocationTargetException;
|
||||
import java.net.MalformedURLException;
|
||||
import java.net.URL;
|
||||
import java.net.URLConnection;
|
||||
import java.net.URLStreamHandler;
|
||||
import java.nio.ByteBuffer;
|
||||
import java.security.CodeSource;
|
||||
import java.security.PermissionCollection;
|
||||
import java.security.ProtectionDomain;
|
||||
import java.security.SecureClassLoader;
|
||||
import java.security.cert.Certificate;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.HashMap;
|
||||
|
|
@ -32,10 +43,22 @@ import java.util.Locale;
|
|||
import java.util.Map;
|
||||
import java.util.StringTokenizer;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.ThreadLocalRandom;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
import com.google.common.io.ByteStreams;
|
||||
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
|
||||
import com.datastax.driver.core.DataType;
|
||||
import org.apache.cassandra.concurrent.JMXEnabledScheduledThreadPoolExecutor;
|
||||
import org.apache.cassandra.concurrent.NamedThreadFactory;
|
||||
import org.apache.cassandra.config.DatabaseDescriptor;
|
||||
import org.apache.cassandra.cql3.ColumnIdentifier;
|
||||
import org.apache.cassandra.db.marshal.AbstractType;
|
||||
import org.apache.cassandra.exceptions.InvalidRequestException;
|
||||
import org.apache.cassandra.utils.FBUtilities;
|
||||
import org.eclipse.jdt.core.compiler.IProblem;
|
||||
import org.eclipse.jdt.internal.compiler.*;
|
||||
|
|
@ -48,27 +71,27 @@ import org.eclipse.jdt.internal.compiler.env.NameEnvironmentAnswer;
|
|||
import org.eclipse.jdt.internal.compiler.impl.CompilerOptions;
|
||||
import org.eclipse.jdt.internal.compiler.problem.DefaultProblemFactory;
|
||||
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
|
||||
import com.datastax.driver.core.DataType;
|
||||
import org.apache.cassandra.cql3.ColumnIdentifier;
|
||||
import org.apache.cassandra.db.marshal.AbstractType;
|
||||
import org.apache.cassandra.exceptions.InvalidRequestException;
|
||||
|
||||
/**
|
||||
* Java source UDF code generation.
|
||||
*/
|
||||
public final class JavaSourceUDFFactory
|
||||
final class JavaBasedUDFunction extends UDFunction
|
||||
{
|
||||
private static final String GENERATED_PACKAGE = "org.apache.cassandra.cql3.udf.gen";
|
||||
private static final String BASE_PACKAGE = "org.apache.cassandra.cql3.udf.gen";
|
||||
|
||||
static final Logger logger = LoggerFactory.getLogger(JavaSourceUDFFactory.class);
|
||||
static final Logger logger = LoggerFactory.getLogger(JavaBasedUDFunction.class);
|
||||
|
||||
private static final AtomicInteger classSequence = new AtomicInteger();
|
||||
|
||||
private static final ClassLoader baseClassLoader = Thread.currentThread().getContextClassLoader();
|
||||
private static final JMXEnabledScheduledThreadPoolExecutor executor =
|
||||
new JMXEnabledScheduledThreadPoolExecutor(
|
||||
DatabaseDescriptor.getMaxHintsThread(),
|
||||
new NamedThreadFactory("UserDefinedFunctions",
|
||||
Thread.MIN_PRIORITY,
|
||||
udfClassLoader,
|
||||
new SecurityThreadGroup("UserDefinedFunctions", null)),
|
||||
"userfunction");
|
||||
|
||||
private static final EcjTargetClassLoader targetClassLoader = new EcjTargetClassLoader();
|
||||
|
||||
private static final ProtectionDomain protectionDomain;
|
||||
|
||||
private static final IErrorHandlingPolicy errorHandlingPolicy = DefaultErrorHandlingPolicies.proceedWithAllProblems();
|
||||
private static final IProblemFactory problemFactory = new DefaultProblemFactory(Locale.ENGLISH);
|
||||
private static final CompilerOptions compilerOptions;
|
||||
|
|
@ -97,7 +120,7 @@ public final class JavaSourceUDFFactory
|
|||
compilerOptions = new CompilerOptions(settings);
|
||||
compilerOptions.parseLiteralExpressionsAsConstants = true;
|
||||
|
||||
try (InputStream input = JavaSourceUDFFactory.class.getResource("JavaSourceUDF.txt").openConnection().getInputStream())
|
||||
try (InputStream input = JavaBasedUDFunction.class.getResource("JavaSourceUDF.txt").openConnection().getInputStream())
|
||||
{
|
||||
ByteArrayOutputStream output = new ByteArrayOutputStream();
|
||||
FBUtilities.copy(input, output, Long.MAX_VALUE);
|
||||
|
|
@ -112,26 +135,42 @@ public final class JavaSourceUDFFactory
|
|||
{
|
||||
throw new RuntimeException(e);
|
||||
}
|
||||
|
||||
CodeSource codeSource;
|
||||
try
|
||||
{
|
||||
codeSource = new CodeSource(new URL("udf", "localhost", 0, "/java", new URLStreamHandler()
|
||||
{
|
||||
protected URLConnection openConnection(URL u)
|
||||
{
|
||||
return null;
|
||||
}
|
||||
}), (Certificate[])null);
|
||||
}
|
||||
catch (MalformedURLException e)
|
||||
{
|
||||
throw new RuntimeException(e);
|
||||
}
|
||||
|
||||
protectionDomain = new ProtectionDomain(codeSource, ThreadAwareSecurityManager.noPermissions, targetClassLoader, null);
|
||||
}
|
||||
|
||||
static UDFunction buildUDF(FunctionName name,
|
||||
List<ColumnIdentifier> argNames,
|
||||
List<AbstractType<?>> argTypes,
|
||||
AbstractType<?> returnType,
|
||||
boolean calledOnNullInput,
|
||||
String body)
|
||||
throws InvalidRequestException
|
||||
private final JavaUDF javaUDF;
|
||||
|
||||
JavaBasedUDFunction(FunctionName name, List<ColumnIdentifier> argNames, List<AbstractType<?>> argTypes,
|
||||
AbstractType<?> returnType, boolean calledOnNullInput, String body)
|
||||
{
|
||||
// argDataTypes is just the C* internal argTypes converted to the Java Driver DataType
|
||||
DataType[] argDataTypes = UDHelper.driverTypes(argTypes);
|
||||
// returnDataType is just the C* internal returnType converted to the Java Driver DataType
|
||||
DataType returnDataType = UDHelper.driverType(returnType);
|
||||
super(name, argNames, argTypes, UDHelper.driverTypes(argTypes),
|
||||
returnType, UDHelper.driverType(returnType), calledOnNullInput, "java", body);
|
||||
|
||||
// javaParamTypes is just the Java representation for argTypes resp. argDataTypes
|
||||
Class<?>[] javaParamTypes = UDHelper.javaTypes(argDataTypes, calledOnNullInput);
|
||||
// javaReturnType is just the Java representation for returnType resp. returnDataType
|
||||
Class<?> javaReturnType = returnDataType.asJavaClass();
|
||||
|
||||
String clsName = generateClassName(name);
|
||||
// put each UDF in a separate package to prevent cross-UDF code access
|
||||
String pkgName = BASE_PACKAGE + '.' + generateClassName(name, 'p');
|
||||
String clsName = generateClassName(name, 'C');
|
||||
|
||||
StringBuilder javaSourceBuilder = new StringBuilder();
|
||||
int lineOffset = 1;
|
||||
|
|
@ -144,6 +183,9 @@ public final class JavaSourceUDFFactory
|
|||
{
|
||||
switch (s)
|
||||
{
|
||||
case "package_name":
|
||||
s = pkgName;
|
||||
break;
|
||||
case "class_name":
|
||||
s = clsName;
|
||||
break;
|
||||
|
|
@ -166,7 +208,7 @@ public final class JavaSourceUDFFactory
|
|||
javaSourceBuilder.append(s);
|
||||
}
|
||||
|
||||
String targetClassName = GENERATED_PACKAGE + '.' + clsName;
|
||||
String targetClassName = pkgName + '.' + clsName;
|
||||
|
||||
String javaSource = javaSourceBuilder.toString();
|
||||
|
||||
|
|
@ -176,11 +218,11 @@ public final class JavaSourceUDFFactory
|
|||
{
|
||||
EcjCompilationUnit compilationUnit = new EcjCompilationUnit(javaSource, targetClassName);
|
||||
|
||||
Compiler compiler = new Compiler(compilationUnit,
|
||||
errorHandlingPolicy,
|
||||
compilerOptions,
|
||||
compilationUnit,
|
||||
problemFactory);
|
||||
org.eclipse.jdt.internal.compiler.Compiler compiler = new Compiler(compilationUnit,
|
||||
errorHandlingPolicy,
|
||||
compilerOptions,
|
||||
compilationUnit,
|
||||
problemFactory);
|
||||
compiler.compile(new ICompilationUnit[]{ compilationUnit });
|
||||
|
||||
if (compilationUnit.problemList != null && !compilationUnit.problemList.isEmpty())
|
||||
|
|
@ -220,18 +262,26 @@ public final class JavaSourceUDFFactory
|
|||
throw new InvalidRequestException("Java source compilation failed:\n" + problems);
|
||||
}
|
||||
|
||||
Class cls = targetClassLoader.loadClass(targetClassName);
|
||||
Thread thread = Thread.currentThread();
|
||||
ClassLoader orig = thread.getContextClassLoader();
|
||||
try
|
||||
{
|
||||
thread.setContextClassLoader(UDFunction.udfClassLoader);
|
||||
// Execute UDF intiialization from UDF class loader
|
||||
|
||||
if (cls.getDeclaredMethods().length != 2 || cls.getDeclaredConstructors().length != 1)
|
||||
throw new InvalidRequestException("Check your source to not define additional Java methods or constructors");
|
||||
MethodType methodType = MethodType.methodType(void.class)
|
||||
.appendParameterTypes(FunctionName.class, List.class, List.class, DataType[].class,
|
||||
AbstractType.class, DataType.class,
|
||||
boolean.class, String.class);
|
||||
MethodHandle ctor = MethodHandles.lookup().findConstructor(cls, methodType);
|
||||
return (UDFunction) ctor.invokeWithArguments(name, argNames, argTypes, argDataTypes,
|
||||
returnType, returnDataType,
|
||||
calledOnNullInput, body);
|
||||
Class cls = targetClassLoader.loadClass(targetClassName);
|
||||
|
||||
if (cls.getDeclaredMethods().length != 2 || cls.getDeclaredConstructors().length != 1)
|
||||
throw new InvalidRequestException("Check your source to not define additional Java methods or constructors");
|
||||
MethodType methodType = MethodType.methodType(void.class)
|
||||
.appendParameterTypes(DataType.class, DataType[].class);
|
||||
MethodHandle ctor = MethodHandles.lookup().findConstructor(cls, methodType);
|
||||
this.javaUDF = (JavaUDF) ctor.invokeWithArguments(returnDataType, argDataTypes);
|
||||
}
|
||||
finally
|
||||
{
|
||||
thread.setContextClassLoader(orig);
|
||||
}
|
||||
}
|
||||
catch (InvocationTargetException e)
|
||||
{
|
||||
|
|
@ -248,6 +298,17 @@ public final class JavaSourceUDFFactory
|
|||
}
|
||||
}
|
||||
|
||||
protected ExecutorService executor()
|
||||
{
|
||||
return executor;
|
||||
}
|
||||
|
||||
protected ByteBuffer executeUserDefined(int protocolVersion, List<ByteBuffer> params)
|
||||
{
|
||||
return javaUDF.executeImpl(protocolVersion, params);
|
||||
}
|
||||
|
||||
|
||||
private static int countNewlines(StringBuilder javaSource)
|
||||
{
|
||||
int ln = 0;
|
||||
|
|
@ -257,19 +318,23 @@ public final class JavaSourceUDFFactory
|
|||
return ln;
|
||||
}
|
||||
|
||||
private static String generateClassName(FunctionName name)
|
||||
private static String generateClassName(FunctionName name, char prefix)
|
||||
{
|
||||
String qualifiedName = name.toString();
|
||||
|
||||
StringBuilder sb = new StringBuilder(qualifiedName.length() + 10);
|
||||
sb.append('C');
|
||||
sb.append(prefix);
|
||||
for (int i = 0; i < qualifiedName.length(); i++)
|
||||
{
|
||||
char c = qualifiedName.charAt(i);
|
||||
if (Character.isJavaIdentifierPart(c))
|
||||
sb.append(c);
|
||||
else
|
||||
sb.append(Integer.toHexString(((short)c)&0xffff));
|
||||
}
|
||||
sb.append('_')
|
||||
.append(ThreadLocalRandom.current().nextInt() & 0xffffff)
|
||||
.append('_')
|
||||
.append(classSequence.incrementAndGet());
|
||||
return sb.toString();
|
||||
}
|
||||
|
|
@ -304,11 +369,11 @@ public final class JavaSourceUDFFactory
|
|||
code.append(",\n");
|
||||
|
||||
if (logger.isDebugEnabled())
|
||||
code.append(" /* parameter '").append(argNames.get(i)).append("' */\n");
|
||||
code.append(" /* parameter '").append(argNames.get(i)).append("' */\n");
|
||||
|
||||
code
|
||||
// cast to Java type
|
||||
.append(" (").append(javaSourceName(paramTypes[i])).append(") ")
|
||||
.append(" (").append(javaSourceName(paramTypes[i])).append(") ")
|
||||
// generate object representation of input parameter (call UDFunction.compose)
|
||||
.append(composeMethod(paramTypes[i])).append("(protocolVersion, ").append(i).append(", params.get(").append(i).append("))");
|
||||
}
|
||||
|
|
@ -432,7 +497,7 @@ public final class JavaSourceUDFFactory
|
|||
|
||||
String resourceName = className.replace('.', '/') + ".class";
|
||||
|
||||
try (InputStream is = baseClassLoader.getResourceAsStream(resourceName))
|
||||
try (InputStream is = UDFunction.udfClassLoader.getResourceAsStream(resourceName))
|
||||
{
|
||||
if (is != null)
|
||||
{
|
||||
|
|
@ -454,7 +519,7 @@ public final class JavaSourceUDFFactory
|
|||
if (result.equals(this.className))
|
||||
return false;
|
||||
String resourceName = result.replace('.', '/') + ".class";
|
||||
try (InputStream is = baseClassLoader.getResourceAsStream(resourceName))
|
||||
try (InputStream is = UDFunction.udfClassLoader.getResourceAsStream(resourceName))
|
||||
{
|
||||
return is == null;
|
||||
}
|
||||
|
|
@ -493,8 +558,13 @@ public final class JavaSourceUDFFactory
|
|||
}
|
||||
}
|
||||
|
||||
static final class EcjTargetClassLoader extends ClassLoader
|
||||
static final class EcjTargetClassLoader extends SecureClassLoader
|
||||
{
|
||||
EcjTargetClassLoader()
|
||||
{
|
||||
super(UDFunction.udfClassLoader);
|
||||
}
|
||||
|
||||
// This map is usually empty.
|
||||
// It only contains data *during* UDF compilation but not during runtime.
|
||||
//
|
||||
|
|
@ -503,12 +573,7 @@ public final class JavaSourceUDFFactory
|
|||
//
|
||||
private final Map<String, byte[]> classes = new ConcurrentHashMap<>();
|
||||
|
||||
EcjTargetClassLoader()
|
||||
{
|
||||
super(baseClassLoader);
|
||||
}
|
||||
|
||||
public void addClass(String className, byte[] classData)
|
||||
void addClass(String className, byte[] classData)
|
||||
{
|
||||
classes.put(className, classData);
|
||||
}
|
||||
|
|
@ -518,8 +583,14 @@ public final class JavaSourceUDFFactory
|
|||
// remove the class binary - it's only used once - so it's wasting heap
|
||||
byte[] classData = classes.remove(name);
|
||||
|
||||
return classData != null ? defineClass(name, classData, 0, classData.length)
|
||||
: super.findClass(name);
|
||||
if (classData != null)
|
||||
return defineClass(name, classData, 0, classData.length, protectionDomain);
|
||||
|
||||
return getParent().loadClass(name);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
protected PermissionCollection getPermissions(CodeSource codesource)
|
||||
{
|
||||
return ThreadAwareSecurityManager.noPermissions;
|
||||
}
|
||||
}}
|
||||
|
|
@ -0,0 +1,105 @@
|
|||
/*
|
||||
* Licensed to the Apache Software Foundation (ASF) under one
|
||||
* or more contributor license agreements. See the NOTICE file
|
||||
* distributed with this work for additional information
|
||||
* regarding copyright ownership. The ASF licenses this file
|
||||
* to you under the Apache License, Version 2.0 (the
|
||||
* "License"); you may not use this file except in compliance
|
||||
* with the License. You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
package org.apache.cassandra.cql3.functions;
|
||||
|
||||
import java.nio.ByteBuffer;
|
||||
import java.util.List;
|
||||
|
||||
import com.datastax.driver.core.DataType;
|
||||
import com.datastax.driver.core.ProtocolVersion;
|
||||
import org.apache.cassandra.db.marshal.AbstractType;
|
||||
|
||||
/**
|
||||
* Base class for all Java UDFs.
|
||||
* Used to separate internal classes like {@link UDFunction} from user provided code.
|
||||
* Only references <b>to</b> this class (and generated implementations) are allowed -
|
||||
* references from this class back to C* code are not allowed (except argument/return type information).
|
||||
*/
|
||||
public abstract class JavaUDF
|
||||
{
|
||||
private final DataType returnDataType;
|
||||
private final DataType[] argDataTypes;
|
||||
|
||||
protected JavaUDF(DataType returnDataType, DataType[] argDataTypes)
|
||||
{
|
||||
this.returnDataType = returnDataType;
|
||||
this.argDataTypes = argDataTypes;
|
||||
}
|
||||
|
||||
protected abstract ByteBuffer executeImpl(int protocolVersion, List<ByteBuffer> params);
|
||||
|
||||
protected Object compose(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
return UDFunction.compose(argDataTypes, protocolVersion, argIndex, value);
|
||||
}
|
||||
|
||||
protected ByteBuffer decompose(int protocolVersion, Object value)
|
||||
{
|
||||
return UDFunction.decompose(returnDataType, protocolVersion, value);
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected float compose_float(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (float) DataType.cfloat().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected double compose_double(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (double) DataType.cdouble().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected byte compose_byte(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (byte) DataType.tinyint().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected short compose_short(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (short) DataType.smallint().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected int compose_int(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (int) DataType.cint().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected long compose_long(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (long) DataType.bigint().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected boolean compose_boolean(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (boolean) DataType.cboolean().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
}
|
||||
|
|
@ -1,150 +0,0 @@
|
|||
/*
|
||||
* Licensed to the Apache Software Foundation (ASF) under one
|
||||
* or more contributor license agreements. See the NOTICE file
|
||||
* distributed with this work for additional information
|
||||
* regarding copyright ownership. The ASF licenses this file
|
||||
* to you under the Apache License, Version 2.0 (the
|
||||
* "License"); you may not use this file except in compliance
|
||||
* with the License. You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
package org.apache.cassandra.cql3.functions;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.math.BigInteger;
|
||||
import java.nio.ByteBuffer;
|
||||
import java.util.HashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import javax.script.Bindings;
|
||||
import javax.script.Compilable;
|
||||
import javax.script.CompiledScript;
|
||||
import javax.script.ScriptEngine;
|
||||
import javax.script.ScriptEngineFactory;
|
||||
import javax.script.ScriptEngineManager;
|
||||
import javax.script.ScriptException;
|
||||
import javax.script.SimpleBindings;
|
||||
|
||||
import org.apache.cassandra.cql3.ColumnIdentifier;
|
||||
import org.apache.cassandra.db.marshal.AbstractType;
|
||||
import org.apache.cassandra.exceptions.FunctionExecutionException;
|
||||
import org.apache.cassandra.exceptions.InvalidRequestException;
|
||||
|
||||
public class ScriptBasedUDF extends UDFunction
|
||||
{
|
||||
static final Map<String, Compilable> scriptEngines = new HashMap<>();
|
||||
|
||||
static {
|
||||
ScriptEngineManager scriptEngineManager = new ScriptEngineManager();
|
||||
for (ScriptEngineFactory scriptEngineFactory : scriptEngineManager.getEngineFactories())
|
||||
{
|
||||
ScriptEngine scriptEngine = scriptEngineFactory.getScriptEngine();
|
||||
boolean compilable = scriptEngine instanceof Compilable;
|
||||
if (compilable)
|
||||
{
|
||||
logger.info("Found scripting engine {} {} - {} {} - language names: {}",
|
||||
scriptEngineFactory.getEngineName(), scriptEngineFactory.getEngineVersion(),
|
||||
scriptEngineFactory.getLanguageName(), scriptEngineFactory.getLanguageVersion(),
|
||||
scriptEngineFactory.getNames());
|
||||
for (String name : scriptEngineFactory.getNames())
|
||||
scriptEngines.put(name, (Compilable) scriptEngine);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private final CompiledScript script;
|
||||
|
||||
ScriptBasedUDF(FunctionName name,
|
||||
List<ColumnIdentifier> argNames,
|
||||
List<AbstractType<?>> argTypes,
|
||||
AbstractType<?> returnType,
|
||||
boolean calledOnNullInput,
|
||||
String language,
|
||||
String body)
|
||||
throws InvalidRequestException
|
||||
{
|
||||
super(name, argNames, argTypes, returnType, calledOnNullInput, language, body);
|
||||
|
||||
Compilable scriptEngine = scriptEngines.get(language);
|
||||
if (scriptEngine == null)
|
||||
throw new InvalidRequestException(String.format("Invalid language '%s' for function '%s'", language, name));
|
||||
|
||||
try
|
||||
{
|
||||
this.script = scriptEngine.compile(body);
|
||||
}
|
||||
catch (RuntimeException | ScriptException e)
|
||||
{
|
||||
logger.info("Failed to compile function '{}' for language {}: ", name, language, e);
|
||||
throw new InvalidRequestException(
|
||||
String.format("Failed to compile function '%s' for language %s: %s", name, language, e));
|
||||
}
|
||||
}
|
||||
|
||||
public ByteBuffer executeUserDefined(int protocolVersion, List<ByteBuffer> parameters) throws InvalidRequestException
|
||||
{
|
||||
Object[] params = new Object[argTypes.size()];
|
||||
for (int i = 0; i < params.length; i++)
|
||||
params[i] = compose(protocolVersion, i, parameters.get(i));
|
||||
|
||||
try
|
||||
{
|
||||
Bindings bindings = new SimpleBindings();
|
||||
for (int i = 0; i < params.length; i++)
|
||||
bindings.put(argNames.get(i).toString(), params[i]);
|
||||
|
||||
Object result = script.eval(bindings);
|
||||
if (result == null)
|
||||
return null;
|
||||
|
||||
Class<?> javaReturnType = returnDataType.asJavaClass();
|
||||
Class<?> resultType = result.getClass();
|
||||
if (!javaReturnType.isAssignableFrom(resultType))
|
||||
{
|
||||
if (result instanceof Number)
|
||||
{
|
||||
Number rNumber = (Number) result;
|
||||
if (javaReturnType == Integer.class)
|
||||
result = rNumber.intValue();
|
||||
else if (javaReturnType == Short.class)
|
||||
result = rNumber.shortValue();
|
||||
else if (javaReturnType == Byte.class)
|
||||
result = rNumber.byteValue();
|
||||
else if (javaReturnType == Long.class)
|
||||
result = rNumber.longValue();
|
||||
else if (javaReturnType == Float.class)
|
||||
result = rNumber.floatValue();
|
||||
else if (javaReturnType == Double.class)
|
||||
result = rNumber.doubleValue();
|
||||
else if (javaReturnType == BigInteger.class)
|
||||
{
|
||||
if (rNumber instanceof BigDecimal)
|
||||
result = ((BigDecimal)rNumber).toBigInteger();
|
||||
else if (rNumber instanceof Double || rNumber instanceof Float)
|
||||
result = new BigDecimal(rNumber.toString()).toBigInteger();
|
||||
else
|
||||
result = BigInteger.valueOf(rNumber.longValue());
|
||||
}
|
||||
else if (javaReturnType == BigDecimal.class)
|
||||
// String c'tor of BigDecimal is more accurate than valueOf(double)
|
||||
result = new BigDecimal(rNumber.toString());
|
||||
}
|
||||
}
|
||||
|
||||
return decompose(protocolVersion, result);
|
||||
}
|
||||
catch (RuntimeException | ScriptException e)
|
||||
{
|
||||
logger.debug("Execution of UDF '{}' failed", name, e);
|
||||
throw FunctionExecutionException.create(this, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,265 @@
|
|||
/*
|
||||
* Licensed to the Apache Software Foundation (ASF) under one
|
||||
* or more contributor license agreements. See the NOTICE file
|
||||
* distributed with this work for additional information
|
||||
* regarding copyright ownership. The ASF licenses this file
|
||||
* to you under the Apache License, Version 2.0 (the
|
||||
* "License"); you may not use this file except in compliance
|
||||
* with the License. You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
package org.apache.cassandra.cql3.functions;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.math.BigInteger;
|
||||
import java.net.MalformedURLException;
|
||||
import java.net.URL;
|
||||
import java.net.URLConnection;
|
||||
import java.net.URLStreamHandler;
|
||||
import java.nio.ByteBuffer;
|
||||
import java.security.AccessControlContext;
|
||||
import java.security.AccessController;
|
||||
import java.security.CodeSource;
|
||||
import java.security.PrivilegedActionException;
|
||||
import java.security.PrivilegedExceptionAction;
|
||||
import java.security.ProtectionDomain;
|
||||
import java.security.cert.Certificate;
|
||||
import java.util.Arrays;
|
||||
import java.util.Collections;
|
||||
import java.util.HashMap;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import javax.script.Bindings;
|
||||
import javax.script.Compilable;
|
||||
import javax.script.CompiledScript;
|
||||
import javax.script.ScriptContext;
|
||||
import javax.script.ScriptEngine;
|
||||
import javax.script.ScriptEngineFactory;
|
||||
import javax.script.ScriptEngineManager;
|
||||
import javax.script.ScriptException;
|
||||
import javax.script.SimpleScriptContext;
|
||||
|
||||
import org.apache.cassandra.concurrent.JMXEnabledScheduledThreadPoolExecutor;
|
||||
import org.apache.cassandra.concurrent.NamedThreadFactory;
|
||||
import org.apache.cassandra.config.DatabaseDescriptor;
|
||||
import org.apache.cassandra.cql3.ColumnIdentifier;
|
||||
import org.apache.cassandra.db.marshal.AbstractType;
|
||||
import org.apache.cassandra.exceptions.InvalidRequestException;
|
||||
|
||||
final class ScriptBasedUDFunction extends UDFunction
|
||||
{
|
||||
static final Map<String, Compilable> scriptEngines = new HashMap<>();
|
||||
|
||||
private static final ProtectionDomain protectionDomain;
|
||||
private static final AccessControlContext accessControlContext;
|
||||
|
||||
//
|
||||
// For scripted UDFs we have to rely on the security mechanisms of the scripting engine and
|
||||
// SecurityManager - especially SecurityManager.checkPackageAccess(). Unlike Java-UDFs, strict checking
|
||||
// of class access via the UDF class loader is not possible, since e.g. Nashorn builds its own class loader
|
||||
// (jdk.nashorn.internal.runtime.ScriptLoader / jdk.nashorn.internal.runtime.NashornLoader) configured with
|
||||
// a system class loader.
|
||||
//
|
||||
private static final String[] allowedPackagesArray =
|
||||
{
|
||||
// following required by jdk.nashorn.internal.objects.Global.initJavaAccess()
|
||||
"",
|
||||
"com",
|
||||
"edu",
|
||||
"java",
|
||||
"javax",
|
||||
"javafx",
|
||||
"org",
|
||||
// following required by Nashorn runtime
|
||||
"java.lang",
|
||||
"java.lang.invoke",
|
||||
"java.lang.reflect",
|
||||
"java.nio.charset",
|
||||
"java.util",
|
||||
"java.util.concurrent",
|
||||
"javax.script",
|
||||
"sun.reflect",
|
||||
"jdk.internal.org.objectweb.asm.commons",
|
||||
"jdk.nashorn.internal.runtime",
|
||||
"jdk.nashorn.internal.runtime.linker",
|
||||
// following required by Java Driver
|
||||
"java.math",
|
||||
"java.text",
|
||||
"com.google.common.base",
|
||||
"com.google.common.reflect",
|
||||
// following required by UDF
|
||||
"com.datastax.driver.core",
|
||||
"com.datastax.driver.core.utils"
|
||||
};
|
||||
|
||||
private static final JMXEnabledScheduledThreadPoolExecutor executor =
|
||||
new JMXEnabledScheduledThreadPoolExecutor(
|
||||
DatabaseDescriptor.getMaxHintsThread(),
|
||||
new NamedThreadFactory("UserDefinedScriptFunctions",
|
||||
Thread.MIN_PRIORITY,
|
||||
udfClassLoader,
|
||||
new SecurityThreadGroup("UserDefinedScriptFunctions", Collections.unmodifiableSet(new HashSet<>(Arrays.asList(allowedPackagesArray))))),
|
||||
"userscripts");
|
||||
|
||||
static
|
||||
{
|
||||
ScriptEngineManager scriptEngineManager = new ScriptEngineManager();
|
||||
for (ScriptEngineFactory scriptEngineFactory : scriptEngineManager.getEngineFactories())
|
||||
{
|
||||
ScriptEngine scriptEngine = scriptEngineFactory.getScriptEngine();
|
||||
boolean compilable = scriptEngine instanceof Compilable;
|
||||
if (compilable)
|
||||
{
|
||||
logger.info("Found scripting engine {} {} - {} {} - language names: {}",
|
||||
scriptEngineFactory.getEngineName(), scriptEngineFactory.getEngineVersion(),
|
||||
scriptEngineFactory.getLanguageName(), scriptEngineFactory.getLanguageVersion(),
|
||||
scriptEngineFactory.getNames());
|
||||
for (String name : scriptEngineFactory.getNames())
|
||||
scriptEngines.put(name, (Compilable) scriptEngine);
|
||||
}
|
||||
}
|
||||
|
||||
try
|
||||
{
|
||||
protectionDomain = new ProtectionDomain(new CodeSource(new URL("udf", "localhost", 0, "/script", new URLStreamHandler()
|
||||
{
|
||||
protected URLConnection openConnection(URL u)
|
||||
{
|
||||
return null;
|
||||
}
|
||||
}), (Certificate[]) null), ThreadAwareSecurityManager.noPermissions);
|
||||
}
|
||||
catch (MalformedURLException e)
|
||||
{
|
||||
throw new RuntimeException(e);
|
||||
}
|
||||
accessControlContext = new AccessControlContext(new ProtectionDomain[]{ protectionDomain });
|
||||
}
|
||||
|
||||
private final CompiledScript script;
|
||||
|
||||
ScriptBasedUDFunction(FunctionName name,
|
||||
List<ColumnIdentifier> argNames,
|
||||
List<AbstractType<?>> argTypes,
|
||||
AbstractType<?> returnType,
|
||||
boolean calledOnNullInput,
|
||||
String language,
|
||||
String body)
|
||||
{
|
||||
super(name, argNames, argTypes, returnType, calledOnNullInput, language, body);
|
||||
|
||||
Compilable scriptEngine = scriptEngines.get(language);
|
||||
if (scriptEngine == null)
|
||||
throw new InvalidRequestException(String.format("Invalid language '%s' for function '%s'", language, name));
|
||||
|
||||
// execute compilation with no-permissions to prevent evil code e.g. via "static code blocks" / "class initialization"
|
||||
try
|
||||
{
|
||||
this.script = AccessController.doPrivileged((PrivilegedExceptionAction<CompiledScript>) () -> scriptEngine.compile(body),
|
||||
accessControlContext);
|
||||
}
|
||||
catch (PrivilegedActionException x)
|
||||
{
|
||||
Throwable e = x.getCause();
|
||||
logger.info("Failed to compile function '{}' for language {}: ", name, language, e);
|
||||
throw new InvalidRequestException(
|
||||
String.format("Failed to compile function '%s' for language %s: %s", name, language, e));
|
||||
}
|
||||
}
|
||||
|
||||
protected ExecutorService executor()
|
||||
{
|
||||
return executor;
|
||||
}
|
||||
|
||||
public ByteBuffer executeUserDefined(int protocolVersion, List<ByteBuffer> parameters)
|
||||
{
|
||||
Object[] params = new Object[argTypes.size()];
|
||||
for (int i = 0; i < params.length; i++)
|
||||
params[i] = compose(protocolVersion, i, parameters.get(i));
|
||||
|
||||
ScriptContext scriptContext = new SimpleScriptContext();
|
||||
scriptContext.setAttribute("javax.script.filename", this.name.toString(), ScriptContext.ENGINE_SCOPE);
|
||||
Bindings bindings = scriptContext.getBindings(ScriptContext.ENGINE_SCOPE);
|
||||
for (int i = 0; i < params.length; i++)
|
||||
bindings.put(argNames.get(i).toString(), params[i]);
|
||||
|
||||
Object result;
|
||||
try
|
||||
{
|
||||
// How to prevent Class.forName() _without_ "help" from the script engine ?
|
||||
// NOTE: Nashorn enforces a special permission to allow class-loading, which is not granted - so it's fine.
|
||||
|
||||
result = script.eval(scriptContext);
|
||||
}
|
||||
catch (ScriptException e)
|
||||
{
|
||||
throw new RuntimeException(e);
|
||||
}
|
||||
if (result == null)
|
||||
return null;
|
||||
|
||||
Class<?> javaReturnType = returnDataType.asJavaClass();
|
||||
Class<?> resultType = result.getClass();
|
||||
if (!javaReturnType.isAssignableFrom(resultType))
|
||||
{
|
||||
if (result instanceof Number)
|
||||
{
|
||||
Number rNumber = (Number) result;
|
||||
if (javaReturnType == Integer.class)
|
||||
result = rNumber.intValue();
|
||||
else if (javaReturnType == Long.class)
|
||||
result = rNumber.longValue();
|
||||
else if (javaReturnType == Short.class)
|
||||
result = rNumber.shortValue();
|
||||
else if (javaReturnType == Byte.class)
|
||||
result = rNumber.byteValue();
|
||||
else if (javaReturnType == Float.class)
|
||||
result = rNumber.floatValue();
|
||||
else if (javaReturnType == Double.class)
|
||||
result = rNumber.doubleValue();
|
||||
else if (javaReturnType == BigInteger.class)
|
||||
{
|
||||
if (javaReturnType == Integer.class)
|
||||
result = rNumber.intValue();
|
||||
else if (javaReturnType == Short.class)
|
||||
result = rNumber.shortValue();
|
||||
else if (javaReturnType == Byte.class)
|
||||
result = rNumber.byteValue();
|
||||
else if (javaReturnType == Long.class)
|
||||
result = rNumber.longValue();
|
||||
else if (javaReturnType == Float.class)
|
||||
result = rNumber.floatValue();
|
||||
else if (javaReturnType == Double.class)
|
||||
result = rNumber.doubleValue();
|
||||
else if (javaReturnType == BigInteger.class)
|
||||
{
|
||||
if (rNumber instanceof BigDecimal)
|
||||
result = ((BigDecimal) rNumber).toBigInteger();
|
||||
else if (rNumber instanceof Double || rNumber instanceof Float)
|
||||
result = new BigDecimal(rNumber.toString()).toBigInteger();
|
||||
else
|
||||
result = BigInteger.valueOf(rNumber.longValue());
|
||||
}
|
||||
else if (javaReturnType == BigDecimal.class)
|
||||
// String c'tor of BigDecimal is more accurate than valueOf(double)
|
||||
result = new BigDecimal(rNumber.toString());
|
||||
}
|
||||
else if (javaReturnType == BigDecimal.class)
|
||||
// String c'tor of BigDecimal is more accurate than valueOf(double)
|
||||
result = new BigDecimal(rNumber.toString());
|
||||
}
|
||||
}
|
||||
|
||||
return decompose(protocolVersion, result);
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,40 @@
|
|||
/*
|
||||
* Licensed to the Apache Software Foundation (ASF) under one
|
||||
* or more contributor license agreements. See the NOTICE file
|
||||
* distributed with this work for additional information
|
||||
* regarding copyright ownership. The ASF licenses this file
|
||||
* to you under the Apache License, Version 2.0 (the
|
||||
* "License"); you may not use this file except in compliance
|
||||
* with the License. You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
package org.apache.cassandra.cql3.functions;
|
||||
|
||||
import java.util.Set;
|
||||
|
||||
/**
|
||||
* Used by {@link ThreadAwareSecurityManager} to determine whether access-control checks needs to be performed.
|
||||
*/
|
||||
public final class SecurityThreadGroup extends ThreadGroup
|
||||
{
|
||||
private final Set<String> allowedPackages;
|
||||
|
||||
public SecurityThreadGroup(String name, Set<String> allowedPackages)
|
||||
{
|
||||
super(name);
|
||||
this.allowedPackages = allowedPackages;
|
||||
}
|
||||
|
||||
public Set<String> getAllowedPackages()
|
||||
{
|
||||
return allowedPackages;
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,184 @@
|
|||
/*
|
||||
* Licensed to the Apache Software Foundation (ASF) under one
|
||||
* or more contributor license agreements. See the NOTICE file
|
||||
* distributed with this work for additional information
|
||||
* regarding copyright ownership. The ASF licenses this file
|
||||
* to you under the Apache License, Version 2.0 (the
|
||||
* "License"); you may not use this file except in compliance
|
||||
* with the License. You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
package org.apache.cassandra.cql3.functions;
|
||||
|
||||
import java.security.AccessControlException;
|
||||
import java.security.AllPermission;
|
||||
import java.security.CodeSource;
|
||||
import java.security.Permission;
|
||||
import java.security.PermissionCollection;
|
||||
import java.security.Permissions;
|
||||
import java.security.Policy;
|
||||
import java.security.ProtectionDomain;
|
||||
import java.util.Collections;
|
||||
import java.util.Enumeration;
|
||||
import java.util.Set;
|
||||
|
||||
import sun.security.util.SecurityConstants;
|
||||
|
||||
/**
|
||||
* Custom {@link SecurityManager} and {@link Policy} implementation that only performs access checks
|
||||
* if explicitly enabled.
|
||||
* <p>
|
||||
* This implementation gives no measurable performance panalty
|
||||
* (see <a href="http://cstar.datastax.com/tests/id/1d461628-12ba-11e5-918f-42010af0688f">see cstar test</a>).
|
||||
* This is better than the penalty of 1 to 3 percent using a standard {@code SecurityManager} with an <i>allow all</i> policy.
|
||||
* </p>
|
||||
*/
|
||||
public final class ThreadAwareSecurityManager extends SecurityManager
|
||||
{
|
||||
static final PermissionCollection noPermissions = new PermissionCollection()
|
||||
{
|
||||
public void add(Permission permission)
|
||||
{
|
||||
throw new UnsupportedOperationException();
|
||||
}
|
||||
|
||||
public boolean implies(Permission permission)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
public Enumeration<Permission> elements()
|
||||
{
|
||||
return Collections.emptyEnumeration();
|
||||
}
|
||||
};
|
||||
|
||||
private static volatile boolean installed;
|
||||
|
||||
public static void install()
|
||||
{
|
||||
if (installed)
|
||||
return;
|
||||
System.setSecurityManager(new ThreadAwareSecurityManager());
|
||||
installed = true;
|
||||
}
|
||||
|
||||
static
|
||||
{
|
||||
//
|
||||
// Use own security policy to be easier (and faster) since the C* has no fine grained permissions.
|
||||
// Either code has access to everything or code has access to nothing (UDFs).
|
||||
// This also removes the burden to maintain and configure policy files for production, unit tests etc.
|
||||
//
|
||||
// Note: a permission is only granted, if there is no objector. This means that
|
||||
// AccessController/AccessControlContext collect all applicable ProtectionDomains - only if none of these
|
||||
// applicable ProtectionDomains denies access, the permission is granted.
|
||||
// A ProtectionDomain can have its origin at an oridinary code-source or provided via a
|
||||
// AccessController.doPrivileded() call.
|
||||
//
|
||||
Policy.setPolicy(new Policy()
|
||||
{
|
||||
public PermissionCollection getPermissions(CodeSource codesource)
|
||||
{
|
||||
// contract of getPermissions() methods is to return a _mutable_ PermissionCollection
|
||||
|
||||
Permissions perms = new Permissions();
|
||||
|
||||
if (codesource == null || codesource.getLocation() == null)
|
||||
return perms;
|
||||
|
||||
switch (codesource.getLocation().getProtocol())
|
||||
{
|
||||
case "file":
|
||||
// All JARs and class files reside on the file system - we can safely
|
||||
// assume that these classes are "good".
|
||||
perms.add(new AllPermission());
|
||||
return perms;
|
||||
}
|
||||
|
||||
return perms;
|
||||
}
|
||||
|
||||
public PermissionCollection getPermissions(ProtectionDomain domain)
|
||||
{
|
||||
return getPermissions(domain.getCodeSource());
|
||||
}
|
||||
|
||||
public boolean implies(ProtectionDomain domain, Permission permission)
|
||||
{
|
||||
CodeSource codesource = domain.getCodeSource();
|
||||
if (codesource == null || codesource.getLocation() == null)
|
||||
return false;
|
||||
|
||||
switch (codesource.getLocation().getProtocol())
|
||||
{
|
||||
case "file":
|
||||
// All JARs and class files reside on the file system - we can safely
|
||||
// assume that these classes are "good".
|
||||
return true;
|
||||
}
|
||||
|
||||
return false;
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
private ThreadAwareSecurityManager()
|
||||
{
|
||||
}
|
||||
|
||||
private static boolean isSecuredThread()
|
||||
{
|
||||
return Thread.currentThread().getThreadGroup() instanceof SecurityThreadGroup;
|
||||
}
|
||||
|
||||
public void checkAccess(Thread t)
|
||||
{
|
||||
// need to override since the default implementation is kind of ...
|
||||
|
||||
if (isSecuredThread())
|
||||
throw new AccessControlException("access denied: " + SecurityConstants.MODIFY_THREAD_PERMISSION, SecurityConstants.MODIFY_THREAD_PERMISSION);
|
||||
super.checkAccess(t);
|
||||
}
|
||||
|
||||
public void checkAccess(ThreadGroup g)
|
||||
{
|
||||
// need to override since the default implementation is kind of ...
|
||||
|
||||
if (isSecuredThread())
|
||||
throw new AccessControlException("access denied: " + SecurityConstants.MODIFY_THREADGROUP_PERMISSION, SecurityConstants.MODIFY_THREADGROUP_PERMISSION);
|
||||
super.checkAccess(g);
|
||||
}
|
||||
|
||||
public void checkPermission(Permission perm)
|
||||
{
|
||||
if (isSecuredThread())
|
||||
super.checkPermission(perm);
|
||||
}
|
||||
|
||||
public void checkPermission(Permission perm, Object context)
|
||||
{
|
||||
if (isSecuredThread())
|
||||
super.checkPermission(perm, context);
|
||||
}
|
||||
|
||||
public void checkPackageAccess(String pkg)
|
||||
{
|
||||
ThreadGroup threadGroup = Thread.currentThread().getThreadGroup();
|
||||
if (threadGroup instanceof SecurityThreadGroup)
|
||||
{
|
||||
Set<String> allowedPackages = ((SecurityThreadGroup) threadGroup).getAllowedPackages();
|
||||
if (allowedPackages != null && !allowedPackages.contains(pkg))
|
||||
throw new AccessControlException("access denied: " + new RuntimePermission("accessClassInPackage." + pkg), new RuntimePermission("accessClassInPackage." + pkg));
|
||||
super.checkPackageAccess(pkg);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -164,7 +164,7 @@ public class UDAggregate extends AbstractFunction implements AggregateFunction
|
|||
{
|
||||
UDFunction udf = (UDFunction)stateFunction;
|
||||
if (udf.isCallableWrtNullable(fArgs))
|
||||
state = udf.executeUserDefined(protocolVersion, fArgs);
|
||||
state = udf.execute(protocolVersion, fArgs);
|
||||
}
|
||||
else
|
||||
{
|
||||
|
|
|
|||
|
|
@ -17,8 +17,22 @@
|
|||
*/
|
||||
package org.apache.cassandra.cql3.functions;
|
||||
|
||||
import java.lang.management.ManagementFactory;
|
||||
import java.lang.management.ThreadMXBean;
|
||||
import java.net.InetAddress;
|
||||
import java.net.URL;
|
||||
import java.nio.ByteBuffer;
|
||||
import java.util.*;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.Enumeration;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.ExecutionException;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.Future;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.TimeoutException;
|
||||
|
||||
import com.google.common.base.Objects;
|
||||
import org.slf4j.Logger;
|
||||
|
|
@ -27,16 +41,20 @@ import org.slf4j.LoggerFactory;
|
|||
import com.datastax.driver.core.DataType;
|
||||
import com.datastax.driver.core.ProtocolVersion;
|
||||
import com.datastax.driver.core.UserType;
|
||||
import org.apache.cassandra.config.Config;
|
||||
import org.apache.cassandra.config.DatabaseDescriptor;
|
||||
import org.apache.cassandra.config.Schema;
|
||||
import org.apache.cassandra.cql3.*;
|
||||
import org.apache.cassandra.cql3.ColumnIdentifier;
|
||||
import org.apache.cassandra.db.marshal.AbstractType;
|
||||
import org.apache.cassandra.exceptions.*;
|
||||
import org.apache.cassandra.exceptions.FunctionExecutionException;
|
||||
import org.apache.cassandra.exceptions.InvalidRequestException;
|
||||
import org.apache.cassandra.schema.Functions;
|
||||
import org.apache.cassandra.schema.KeyspaceMetadata;
|
||||
import org.apache.cassandra.service.ClientWarn;
|
||||
import org.apache.cassandra.service.MigrationManager;
|
||||
import org.apache.cassandra.tracing.Tracing;
|
||||
import org.apache.cassandra.utils.ByteBufferUtil;
|
||||
import org.apache.cassandra.utils.JVMStabilityInspector;
|
||||
|
||||
/**
|
||||
* Base class for User Defined Functions.
|
||||
|
|
@ -45,6 +63,8 @@ public abstract class UDFunction extends AbstractFunction implements ScalarFunct
|
|||
{
|
||||
protected static final Logger logger = LoggerFactory.getLogger(UDFunction.class);
|
||||
|
||||
static final ThreadMXBean threadMXBean = ManagementFactory.getThreadMXBean();
|
||||
|
||||
protected final List<ColumnIdentifier> argNames;
|
||||
|
||||
protected final String language;
|
||||
|
|
@ -54,6 +74,100 @@ public abstract class UDFunction extends AbstractFunction implements ScalarFunct
|
|||
protected final DataType returnDataType;
|
||||
protected final boolean calledOnNullInput;
|
||||
|
||||
//
|
||||
// Access to classes is controlled via a whitelist and a blacklist.
|
||||
//
|
||||
// When a class is requested (both during compilation and runtime),
|
||||
// the whitelistedPatterns array is searched first, whether the
|
||||
// requested name matches one of the patterns. If not, nothing is
|
||||
// returned from the class-loader - meaning ClassNotFoundException
|
||||
// during runtime and "type could not resolved" during compilation.
|
||||
//
|
||||
// If a whitelisted pattern has been found, the blacklistedPatterns
|
||||
// array is searched for a match. If a match is found, class-loader
|
||||
// rejects access. Otherwise the class/resource can be loaded.
|
||||
//
|
||||
private static final String[] whitelistedPatterns =
|
||||
{
|
||||
"com/datastax/driver/core/",
|
||||
"com/google/common/reflect/TypeToken",
|
||||
"java/io/IOException.class",
|
||||
"java/io/Serializable.class",
|
||||
"java/lang/",
|
||||
"java/math/",
|
||||
"java/nio/Buffer.class",
|
||||
"java/nio/ByteBuffer.class",
|
||||
"java/text/",
|
||||
"java/time/",
|
||||
"java/util/",
|
||||
"org/apache/cassandra/cql3/functions/JavaUDF.class",
|
||||
"org/apache/cassandra/exceptions/",
|
||||
};
|
||||
// Only need to blacklist a pattern, if it would otherwise be allowed via whitelistedPatterns
|
||||
private static final String[] blacklistedPatterns =
|
||||
{
|
||||
"com/datastax/driver/core/Cluster.class",
|
||||
"com/datastax/driver/core/Metrics.class",
|
||||
"com/datastax/driver/core/NettyOptions.class",
|
||||
"com/datastax/driver/core/Session.class",
|
||||
"com/datastax/driver/core/Statement.class",
|
||||
"com/datastax/driver/core/TimestampGenerator.class", // indirectly covers ServerSideTimestampGenerator + ThreadLocalMonotonicTimestampGenerator
|
||||
"java/lang/Compiler.class",
|
||||
"java/lang/Package.class",
|
||||
"java/lang/Process.class",
|
||||
"java/lang/ProcessBuilder.class",
|
||||
"java/lang/ProcessEnvironment.class",
|
||||
"java/lang/ProcessImpl.class",
|
||||
"java/lang/Runnable.class",
|
||||
"java/lang/Runtime.class",
|
||||
"java/lang/Shutdown.class",
|
||||
"java/lang/Thread.class",
|
||||
"java/lang/ThreadGroup.class",
|
||||
"java/lang/ThreadLocal.class",
|
||||
"java/lang/instrument/",
|
||||
"java/lang/invoke/",
|
||||
"java/lang/management/",
|
||||
"java/lang/ref/",
|
||||
"java/lang/reflect/",
|
||||
"java/util/ServiceLoader.class",
|
||||
"java/util/Timer.class",
|
||||
"java/util/concurrent/",
|
||||
"java/util/function/",
|
||||
"java/util/jar/",
|
||||
"java/util/logging/",
|
||||
"java/util/prefs/",
|
||||
"java/util/spi/",
|
||||
"java/util/stream/",
|
||||
"java/util/zip/",
|
||||
};
|
||||
|
||||
static boolean secureResource(String resource)
|
||||
{
|
||||
while (resource.startsWith("/"))
|
||||
resource = resource.substring(1);
|
||||
|
||||
for (String white : whitelistedPatterns)
|
||||
if (resource.startsWith(white))
|
||||
{
|
||||
|
||||
// resource is in whitelistedPatterns, let's see if it is not explicityl blacklisted
|
||||
for (String black : blacklistedPatterns)
|
||||
if (resource.startsWith(black))
|
||||
{
|
||||
logger.trace("access denied: resource {}", resource);
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
logger.trace("access denied: resource {}", resource);
|
||||
return false;
|
||||
}
|
||||
|
||||
// setup the UDF class loader with no parent class loader so that we have full control about what class/resource UDF uses
|
||||
static final ClassLoader udfClassLoader = new UDFClassLoader();
|
||||
|
||||
protected UDFunction(FunctionName name,
|
||||
List<ColumnIdentifier> argNames,
|
||||
List<AbstractType<?>> argTypes,
|
||||
|
|
@ -93,15 +207,16 @@ public abstract class UDFunction extends AbstractFunction implements ScalarFunct
|
|||
boolean calledOnNullInput,
|
||||
String language,
|
||||
String body)
|
||||
throws InvalidRequestException
|
||||
{
|
||||
if (!DatabaseDescriptor.enableUserDefinedFunctions())
|
||||
throw new InvalidRequestException("User-defined-functions are disabled in cassandra.yaml - set enable_user_defined_functions=true to enable if you are aware of the security risks");
|
||||
throw new InvalidRequestException("User-defined functions are disabled in cassandra.yaml - set enable_user_defined_functions=true to enable if you are aware of the security risks");
|
||||
|
||||
switch (language)
|
||||
{
|
||||
case "java": return JavaSourceUDFFactory.buildUDF(name, argNames, argTypes, returnType, calledOnNullInput, body);
|
||||
default: return new ScriptBasedUDF(name, argNames, argTypes, returnType, calledOnNullInput, language, body);
|
||||
case "java":
|
||||
return new JavaBasedUDFunction(name, argNames, argTypes, returnType, calledOnNullInput, body);
|
||||
default:
|
||||
return new ScriptBasedUDFunction(name, argNames, argTypes, returnType, calledOnNullInput, language, body);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -121,11 +236,16 @@ public abstract class UDFunction extends AbstractFunction implements ScalarFunct
|
|||
boolean calledOnNullInput,
|
||||
String language,
|
||||
String body,
|
||||
final InvalidRequestException reason)
|
||||
InvalidRequestException reason)
|
||||
{
|
||||
return new UDFunction(name, argNames, argTypes, returnType, calledOnNullInput, language, body)
|
||||
{
|
||||
public ByteBuffer executeUserDefined(int protocolVersion, List<ByteBuffer> parameters) throws InvalidRequestException
|
||||
protected ExecutorService executor()
|
||||
{
|
||||
return Executors.newSingleThreadExecutor();
|
||||
}
|
||||
|
||||
public ByteBuffer executeUserDefined(int protocolVersion, List<ByteBuffer> parameters)
|
||||
{
|
||||
throw new InvalidRequestException(String.format("Function '%s' exists but hasn't been loaded successfully "
|
||||
+ "for the following reason: %s. Please see the server log for details",
|
||||
|
|
@ -135,7 +255,7 @@ public abstract class UDFunction extends AbstractFunction implements ScalarFunct
|
|||
};
|
||||
}
|
||||
|
||||
public final ByteBuffer execute(int protocolVersion, List<ByteBuffer> parameters) throws InvalidRequestException
|
||||
public final ByteBuffer execute(int protocolVersion, List<ByteBuffer> parameters)
|
||||
{
|
||||
if (!DatabaseDescriptor.enableUserDefinedFunctions())
|
||||
throw new InvalidRequestException("User-defined-functions are disabled in cassandra.yaml - set enable_user_defined_functions=true to enable if you are aware of the security risks");
|
||||
|
|
@ -144,11 +264,147 @@ public abstract class UDFunction extends AbstractFunction implements ScalarFunct
|
|||
return null;
|
||||
|
||||
long tStart = System.nanoTime();
|
||||
ByteBuffer result = executeUserDefined(protocolVersion, parameters);
|
||||
Tracing.trace("Executed UDF {} in {}\u03bcs", name(), (System.nanoTime() - tStart) / 1000);
|
||||
return result;
|
||||
parameters = makeEmptyParametersNull(parameters);
|
||||
|
||||
try
|
||||
{
|
||||
// Using async UDF execution is expensive (adds about 100us overhead per invocation on a Core-i7 MBPr).
|
||||
ByteBuffer result = DatabaseDescriptor.enableUserDefinedFunctionsThreads()
|
||||
? executeAsync(protocolVersion, parameters)
|
||||
: executeUserDefined(protocolVersion, parameters);
|
||||
Tracing.trace("Executed UDF {} in {}\u03bcs", name(), (System.nanoTime() - tStart) / 1000);
|
||||
return result;
|
||||
}
|
||||
catch (InvalidRequestException e)
|
||||
{
|
||||
throw e;
|
||||
}
|
||||
catch (Throwable t)
|
||||
{
|
||||
logger.debug("Invocation of user-defined function '{}' failed", this, t);
|
||||
if (t instanceof VirtualMachineError)
|
||||
throw (VirtualMachineError) t;
|
||||
throw FunctionExecutionException.create(this, t);
|
||||
}
|
||||
}
|
||||
|
||||
private static final class ThreadIdAndCpuTime
|
||||
{
|
||||
long threadId;
|
||||
long cpuTime;
|
||||
|
||||
ThreadIdAndCpuTime()
|
||||
{
|
||||
// Looks weird?
|
||||
// This call "just" links this class to java.lang.management - otherwise UDFs (script UDFs) might fail due to
|
||||
// java.security.AccessControlException: access denied: ("java.lang.RuntimePermission" "accessClassInPackage.java.lang.management")
|
||||
// because class loading would be deferred until setup() is executed - but setup() is called with
|
||||
// limited privileges.
|
||||
threadMXBean.getCurrentThreadCpuTime();
|
||||
//
|
||||
// Get the TypeCodec stuff in Java Driver initialized.
|
||||
DataType.inet().format(InetAddress.getLoopbackAddress());
|
||||
DataType.list(DataType.ascii()).format(Collections.emptyList());
|
||||
}
|
||||
|
||||
void setup()
|
||||
{
|
||||
this.threadId = Thread.currentThread().getId();
|
||||
this.cpuTime = threadMXBean.getCurrentThreadCpuTime();
|
||||
}
|
||||
}
|
||||
|
||||
private ByteBuffer executeAsync(int protocolVersion, List<ByteBuffer> parameters)
|
||||
{
|
||||
ThreadIdAndCpuTime threadIdAndCpuTime = new ThreadIdAndCpuTime();
|
||||
|
||||
Future<ByteBuffer> future = executor().submit(() -> {
|
||||
threadIdAndCpuTime.setup();
|
||||
return executeUserDefined(protocolVersion, parameters);
|
||||
});
|
||||
|
||||
try
|
||||
{
|
||||
if (DatabaseDescriptor.getUserDefinedFunctionWarnTimeout() > 0)
|
||||
try
|
||||
{
|
||||
return future.get(DatabaseDescriptor.getUserDefinedFunctionWarnTimeout(), TimeUnit.MILLISECONDS);
|
||||
}
|
||||
catch (TimeoutException e)
|
||||
{
|
||||
|
||||
// log and emit a warning that UDF execution took long
|
||||
String warn = String.format("User defined function %s ran longer than %dms", this, DatabaseDescriptor.getUserDefinedFunctionWarnTimeout());
|
||||
logger.warn(warn);
|
||||
ClientWarn.warn(warn);
|
||||
}
|
||||
|
||||
// retry with difference of warn-timeout to fail-timeout
|
||||
return future.get(DatabaseDescriptor.getUserDefinedFunctionFailTimeout() - DatabaseDescriptor.getUserDefinedFunctionWarnTimeout(), TimeUnit.MILLISECONDS);
|
||||
}
|
||||
catch (InterruptedException e)
|
||||
{
|
||||
Thread.currentThread().interrupt();
|
||||
throw new RuntimeException(e);
|
||||
}
|
||||
catch (ExecutionException e)
|
||||
{
|
||||
Throwable c = e.getCause();
|
||||
if (c instanceof RuntimeException)
|
||||
throw (RuntimeException) c;
|
||||
throw new RuntimeException(c);
|
||||
}
|
||||
catch (TimeoutException e)
|
||||
{
|
||||
// retry a last time with the difference of UDF-fail-timeout to consumed CPU time (just in case execution hit a badly timed GC)
|
||||
try
|
||||
{
|
||||
long cpuTimeMillis = threadMXBean.getThreadCpuTime(threadIdAndCpuTime.threadId) - threadIdAndCpuTime.cpuTime;
|
||||
cpuTimeMillis /= 1000000L;
|
||||
|
||||
return future.get(Math.max(DatabaseDescriptor.getUserDefinedFunctionFailTimeout() - cpuTimeMillis, 0L),
|
||||
TimeUnit.MILLISECONDS);
|
||||
}
|
||||
catch (InterruptedException e1)
|
||||
{
|
||||
Thread.currentThread().interrupt();
|
||||
throw new RuntimeException(e);
|
||||
}
|
||||
catch (ExecutionException e1)
|
||||
{
|
||||
Throwable c = e.getCause();
|
||||
if (c instanceof RuntimeException)
|
||||
throw (RuntimeException) c;
|
||||
throw new RuntimeException(c);
|
||||
}
|
||||
catch (TimeoutException e1)
|
||||
{
|
||||
TimeoutException cause = new TimeoutException(String.format("User defined function %s ran longer than %dms%s",
|
||||
this,
|
||||
DatabaseDescriptor.getUserDefinedFunctionFailTimeout(),
|
||||
DatabaseDescriptor.getUserFunctionTimeoutPolicy() == Config.UserFunctionTimeoutPolicy.ignore
|
||||
? "" : " - will stop Cassandra VM"));
|
||||
FunctionExecutionException fe = FunctionExecutionException.create(this, cause);
|
||||
JVMStabilityInspector.userFunctionTimeout(cause);
|
||||
throw fe;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private List<ByteBuffer> makeEmptyParametersNull(List<ByteBuffer> parameters)
|
||||
{
|
||||
List<ByteBuffer> r = new ArrayList<>(parameters.size());
|
||||
for (int i = 0; i < parameters.size(); i++)
|
||||
{
|
||||
ByteBuffer param = parameters.get(i);
|
||||
r.add(UDHelper.isNullOrEmpty(argTypes.get(i), param)
|
||||
? null : param);
|
||||
}
|
||||
return r;
|
||||
}
|
||||
|
||||
protected abstract ExecutorService executor();
|
||||
|
||||
public boolean isCallableWrtNullable(List<ByteBuffer> parameters)
|
||||
{
|
||||
if (!calledOnNullInput)
|
||||
|
|
@ -158,7 +414,7 @@ public abstract class UDFunction extends AbstractFunction implements ScalarFunct
|
|||
return true;
|
||||
}
|
||||
|
||||
protected abstract ByteBuffer executeUserDefined(int protocolVersion, List<ByteBuffer> parameters) throws InvalidRequestException;
|
||||
protected abstract ByteBuffer executeUserDefined(int protocolVersion, List<ByteBuffer> parameters);
|
||||
|
||||
public boolean isAggregate()
|
||||
{
|
||||
|
|
@ -191,77 +447,38 @@ public abstract class UDFunction extends AbstractFunction implements ScalarFunct
|
|||
}
|
||||
|
||||
/**
|
||||
* Used by UDF implementations (both Java code generated by {@link org.apache.cassandra.cql3.functions.JavaSourceUDFFactory}
|
||||
* and script executor {@link org.apache.cassandra.cql3.functions.ScriptBasedUDF}) to convert the C*
|
||||
* Used by UDF implementations (both Java code generated by {@link JavaBasedUDFunction}
|
||||
* and script executor {@link ScriptBasedUDFunction}) to convert the C*
|
||||
* serialized representation to the Java object representation.
|
||||
*
|
||||
* @param protocolVersion the native protocol version used for serialization
|
||||
* @param argIndex index of the UDF input argument
|
||||
* @param argIndex index of the UDF input argument
|
||||
*/
|
||||
protected Object compose(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
return UDHelper.isNullOrEmpty(argTypes.get(argIndex), value) ? null : argDataTypes[argIndex].deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
return compose(argDataTypes, protocolVersion, argIndex, value);
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected float compose_float(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
protected static Object compose(DataType[] argDataTypes, int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (float)DataType.cfloat().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected double compose_double(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (double)DataType.cdouble().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected byte compose_byte(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (byte)DataType.tinyint().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected short compose_short(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (short)DataType.smallint().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected int compose_int(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (int)DataType.cint().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected long compose_long(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (long)DataType.bigint().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
// do not remove - used by generated Java UDFs
|
||||
protected boolean compose_boolean(int protocolVersion, int argIndex, ByteBuffer value)
|
||||
{
|
||||
assert value != null && value.remaining() > 0;
|
||||
return (boolean) DataType.cboolean().deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
return value == null ? null : argDataTypes[argIndex].deserialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
/**
|
||||
* Used by UDF implementations (both Java code generated by {@link org.apache.cassandra.cql3.functions.JavaSourceUDFFactory}
|
||||
* and script executor {@link org.apache.cassandra.cql3.functions.ScriptBasedUDF}) to convert the Java
|
||||
* Used by UDF implementations (both Java code generated by {@link JavaBasedUDFunction}
|
||||
* and script executor {@link ScriptBasedUDFunction}) to convert the Java
|
||||
* object representation for the return value to the C* serialized representation.
|
||||
*
|
||||
* @param protocolVersion the native protocol version used for serialization
|
||||
*/
|
||||
protected ByteBuffer decompose(int protocolVersion, Object value)
|
||||
{
|
||||
return value == null ? null : returnDataType.serialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
return decompose(returnDataType, protocolVersion, value);
|
||||
}
|
||||
|
||||
protected static ByteBuffer decompose(DataType dataType, int protocolVersion, Object value)
|
||||
{
|
||||
return value == null ? null : dataType.serialize(value, ProtocolVersion.fromInt(protocolVersion));
|
||||
}
|
||||
|
||||
@Override
|
||||
|
|
@ -315,4 +532,41 @@ public abstract class UDFunction extends AbstractFunction implements ScalarFunct
|
|||
if (updated)
|
||||
MigrationManager.announceNewFunction(this, true);
|
||||
}
|
||||
|
||||
private static class UDFClassLoader extends ClassLoader
|
||||
{
|
||||
// insecureClassLoader is the C* class loader
|
||||
static final ClassLoader insecureClassLoader = Thread.currentThread().getContextClassLoader();
|
||||
|
||||
public URL getResource(String name)
|
||||
{
|
||||
if (!secureResource(name))
|
||||
return null;
|
||||
return insecureClassLoader.getResource(name);
|
||||
}
|
||||
|
||||
protected URL findResource(String name)
|
||||
{
|
||||
return getResource(name);
|
||||
}
|
||||
|
||||
public Enumeration<URL> getResources(String name)
|
||||
{
|
||||
return Collections.emptyEnumeration();
|
||||
}
|
||||
|
||||
protected Class<?> findClass(String name) throws ClassNotFoundException
|
||||
{
|
||||
if (!secureResource(name.replace('.', '/') + ".class"))
|
||||
throw new ClassNotFoundException(name);
|
||||
return insecureClassLoader.loadClass(name);
|
||||
}
|
||||
|
||||
public Class<?> loadClass(String name) throws ClassNotFoundException
|
||||
{
|
||||
if (!secureResource(name.replace('.', '/') + ".class"))
|
||||
throw new ClassNotFoundException(name);
|
||||
return super.loadClass(name);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -36,16 +36,15 @@ import javax.management.remote.JMXConnectorServer;
|
|||
import javax.management.remote.JMXServiceURL;
|
||||
import javax.management.remote.rmi.RMIConnectorServer;
|
||||
|
||||
import com.addthis.metrics3.reporter.config.ReporterConfig;
|
||||
import com.codahale.metrics.Meter;
|
||||
import com.codahale.metrics.MetricRegistryListener;
|
||||
import com.codahale.metrics.SharedMetricRegistries;
|
||||
import com.google.common.util.concurrent.Uninterruptibles;
|
||||
import org.apache.cassandra.gms.Gossiper;
|
||||
import org.apache.cassandra.metrics.DefaultNameFactory;
|
||||
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
|
||||
import com.addthis.metrics3.reporter.config.ReporterConfig;
|
||||
import org.apache.cassandra.concurrent.*;
|
||||
import org.apache.cassandra.config.CFMetaData;
|
||||
import org.apache.cassandra.config.DatabaseDescriptor;
|
||||
|
|
@ -54,12 +53,15 @@ import org.apache.cassandra.db.*;
|
|||
import org.apache.cassandra.db.commitlog.CommitLog;
|
||||
import org.apache.cassandra.exceptions.ConfigurationException;
|
||||
import org.apache.cassandra.exceptions.StartupException;
|
||||
import org.apache.cassandra.gms.Gossiper;
|
||||
import org.apache.cassandra.io.FSError;
|
||||
import org.apache.cassandra.io.sstable.CorruptSSTableException;
|
||||
import org.apache.cassandra.io.util.FileUtils;
|
||||
import org.apache.cassandra.metrics.CassandraMetricsRegistry;
|
||||
import org.apache.cassandra.metrics.DefaultNameFactory;
|
||||
import org.apache.cassandra.metrics.StorageMetrics;
|
||||
import org.apache.cassandra.schema.LegacySchemaMigrator;
|
||||
import org.apache.cassandra.cql3.functions.ThreadAwareSecurityManager;
|
||||
import org.apache.cassandra.thrift.ThriftServer;
|
||||
import org.apache.cassandra.tracing.Tracing;
|
||||
import org.apache.cassandra.utils.*;
|
||||
|
|
@ -150,6 +152,8 @@ public class CassandraDaemon
|
|||
if (FBUtilities.isWindows())
|
||||
WindowsFailedSnapshotTracker.deleteOldSnapshots();
|
||||
|
||||
ThreadAwareSecurityManager.install();
|
||||
|
||||
logSystemInfo();
|
||||
|
||||
CLibrary.tryMlockall();
|
||||
|
|
|
|||
|
|
@ -19,11 +19,14 @@ package org.apache.cassandra.utils;
|
|||
|
||||
import java.io.FileNotFoundException;
|
||||
import java.net.SocketException;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicBoolean;
|
||||
|
||||
import com.google.common.annotations.VisibleForTesting;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
|
||||
import org.apache.cassandra.concurrent.ScheduledExecutors;
|
||||
import org.apache.cassandra.config.Config;
|
||||
import org.apache.cassandra.config.DatabaseDescriptor;
|
||||
import org.apache.cassandra.io.FSError;
|
||||
|
|
@ -77,6 +80,23 @@ public final class JVMStabilityInspector
|
|||
killer.killCurrentJVM(t, quiet);
|
||||
}
|
||||
|
||||
public static void userFunctionTimeout(Throwable t)
|
||||
{
|
||||
switch (DatabaseDescriptor.getUserFunctionTimeoutPolicy())
|
||||
{
|
||||
case die:
|
||||
// policy to give 250ms grace time to
|
||||
ScheduledExecutors.nonPeriodicTasks.schedule(() -> killer.killCurrentJVM(t), 250, TimeUnit.MILLISECONDS);
|
||||
break;
|
||||
case die_immediate:
|
||||
killer.killCurrentJVM(t);
|
||||
break;
|
||||
case ignore:
|
||||
logger.error(t.getMessage());
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
@VisibleForTesting
|
||||
public static Killer replaceKiller(Killer newKiller) {
|
||||
Killer oldKiller = JVMStabilityInspector.killer;
|
||||
|
|
@ -87,6 +107,8 @@ public final class JVMStabilityInspector
|
|||
@VisibleForTesting
|
||||
public static class Killer
|
||||
{
|
||||
private final AtomicBoolean killing = new AtomicBoolean();
|
||||
|
||||
/**
|
||||
* Certain situations represent "Die" conditions for the server, and if so, the reason is logged and the current JVM is killed.
|
||||
*
|
||||
|
|
@ -105,8 +127,11 @@ public final class JVMStabilityInspector
|
|||
t.printStackTrace(System.err);
|
||||
logger.error("JVM state determined to be unstable. Exiting forcefully due to:", t);
|
||||
}
|
||||
StorageService.instance.removeShutdownHook();
|
||||
System.exit(100);
|
||||
if (killing.compareAndSet(false, true))
|
||||
{
|
||||
StorageService.instance.removeShutdownHook();
|
||||
System.exit(100);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,39 +1,25 @@
|
|||
package org.apache.cassandra.cql3.udf.gen;
|
||||
package #package_name#;
|
||||
|
||||
import java.nio.ByteBuffer;
|
||||
import java.util.List;
|
||||
import com.datastax.driver.core.DataType;
|
||||
import org.apache.cassandra.cql3.ColumnIdentifier;
|
||||
import org.apache.cassandra.cql3.functions.FunctionName;
|
||||
import org.apache.cassandra.cql3.functions.JavaSourceUDFFactory;
|
||||
import org.apache.cassandra.db.marshal.AbstractType;
|
||||
import org.apache.cassandra.exceptions.FunctionExecutionException;
|
||||
import org.apache.cassandra.exceptions.InvalidRequestException;
|
||||
|
||||
public final class #class_name# extends org.apache.cassandra.cql3.functions.UDFunction
|
||||
import org.apache.cassandra.cql3.functions.JavaUDF;
|
||||
|
||||
import com.datastax.driver.core.DataType;
|
||||
|
||||
public final class #class_name# extends JavaUDF
|
||||
{
|
||||
public #class_name#(FunctionName name, List<ColumnIdentifier> argNames, List<AbstractType<?>> argTypes,
|
||||
DataType[] argDataTypes, AbstractType<?> returnType, DataType returnDataType, boolean calledOnNullInput, String body)
|
||||
public #class_name#(DataType returnDataType, DataType[] argDataTypes)
|
||||
{
|
||||
super(name, argNames, argTypes, argDataTypes, returnType, returnDataType, calledOnNullInput, "java", body);
|
||||
super(returnDataType, argDataTypes);
|
||||
}
|
||||
|
||||
protected ByteBuffer executeUserDefined(int protocolVersion, List<ByteBuffer> params) throws InvalidRequestException
|
||||
public ByteBuffer executeImpl(int protocolVersion, List<ByteBuffer> params)
|
||||
{
|
||||
try
|
||||
{
|
||||
#return_type# result = executeInternal(
|
||||
#return_type# result = executeInternal(
|
||||
#arguments#
|
||||
);
|
||||
return decompose(protocolVersion, result);
|
||||
}
|
||||
catch (Throwable t)
|
||||
{
|
||||
logger.debug("Invocation of function '{}' failed", this, t);
|
||||
if (t instanceof VirtualMachineError)
|
||||
throw (VirtualMachineError)t;
|
||||
throw FunctionExecutionException.create(this, t);
|
||||
}
|
||||
);
|
||||
return decompose(protocolVersion, result);
|
||||
}
|
||||
|
||||
private #return_type# executeInternal(#argument_list#)
|
||||
|
|
|
|||
|
|
@ -23,8 +23,6 @@ import java.util.*;
|
|||
|
||||
import org.junit.After;
|
||||
import org.junit.BeforeClass;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
|
||||
import org.apache.cassandra.cache.CachingOptions;
|
||||
import org.apache.cassandra.config.*;
|
||||
|
|
@ -50,8 +48,6 @@ import org.apache.cassandra.utils.FBUtilities;
|
|||
|
||||
public class SchemaLoader
|
||||
{
|
||||
private static Logger logger = LoggerFactory.getLogger(SchemaLoader.class);
|
||||
|
||||
@BeforeClass
|
||||
public static void loadSchema() throws ConfigurationException
|
||||
{
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ import org.apache.cassandra.config.CFMetaData;
|
|||
import org.apache.cassandra.config.DatabaseDescriptor;
|
||||
import org.apache.cassandra.config.Schema;
|
||||
import org.apache.cassandra.cql3.functions.FunctionName;
|
||||
import org.apache.cassandra.cql3.functions.ThreadAwareSecurityManager;
|
||||
import org.apache.cassandra.cql3.statements.ParsedStatement;
|
||||
import org.apache.cassandra.db.*;
|
||||
import org.apache.cassandra.db.commitlog.CommitLog;
|
||||
|
|
@ -159,6 +160,8 @@ public abstract class CQLTester
|
|||
}
|
||||
});
|
||||
|
||||
ThreadAwareSecurityManager.install();
|
||||
|
||||
Keyspace.setInitialized();
|
||||
isServerPrepared = true;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,478 @@
|
|||
/*
|
||||
* Licensed to the Apache Software Foundation (ASF) under one
|
||||
* or more contributor license agreements. See the NOTICE file
|
||||
* distributed with this work for additional information
|
||||
* regarding copyright ownership. The ASF licenses this file
|
||||
* to you under the Apache License, Version 2.0 (the
|
||||
* "License"); you may not use this file except in compliance
|
||||
* with the License. You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
package org.apache.cassandra.cql3.validation.entities;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.math.BigInteger;
|
||||
import java.util.Arrays;
|
||||
import java.util.Date;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
import java.util.TreeMap;
|
||||
import java.util.TreeSet;
|
||||
import java.util.UUID;
|
||||
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import com.datastax.driver.core.DataType;
|
||||
import com.datastax.driver.core.TupleType;
|
||||
import com.datastax.driver.core.TupleValue;
|
||||
import org.apache.cassandra.cql3.CQLTester;
|
||||
import org.apache.cassandra.cql3.UntypedResultSet;
|
||||
import org.apache.cassandra.cql3.functions.FunctionName;
|
||||
import org.apache.cassandra.exceptions.FunctionExecutionException;
|
||||
import org.apache.cassandra.transport.Server;
|
||||
import org.apache.cassandra.utils.UUIDGen;
|
||||
|
||||
public class UFPureScriptTest extends CQLTester
|
||||
{
|
||||
// Just JavaScript UDFs to check how UDF - especially security/class-loading/sandboxing stuff -
|
||||
// behaves, if no Java UDF has been executed before.
|
||||
|
||||
// Do not add any other test here - especially none using Java UDFs
|
||||
|
||||
@Test
|
||||
public void testJavascriptSimpleCollections() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, lst list<double>, st set<text>, mp map<int, boolean>)");
|
||||
|
||||
String fName1 = createFunction(KEYSPACE_PER_TEST, "list<double>",
|
||||
"CREATE FUNCTION %s( lst list<double> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS list<double> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS 'lst;';");
|
||||
String fName2 = createFunction(KEYSPACE_PER_TEST, "set<text>",
|
||||
"CREATE FUNCTION %s( st set<text> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS set<text> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS 'st;';");
|
||||
String fName3 = createFunction(KEYSPACE_PER_TEST, "map<int, boolean>",
|
||||
"CREATE FUNCTION %s( mp map<int, boolean> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS map<int, boolean> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS 'mp;';");
|
||||
|
||||
List<Double> list = Arrays.asList(1d, 2d, 3d);
|
||||
Set<String> set = new TreeSet<>(Arrays.asList("one", "three", "two"));
|
||||
Map<Integer, Boolean> map = new TreeMap<>();
|
||||
map.put(1, true);
|
||||
map.put(2, false);
|
||||
map.put(3, true);
|
||||
|
||||
execute("INSERT INTO %s (key, lst, st, mp) VALUES (1, ?, ?, ?)", list, set, map);
|
||||
|
||||
assertRows(execute("SELECT lst, st, mp FROM %s WHERE key = 1"),
|
||||
row(list, set, map));
|
||||
|
||||
assertRows(execute("SELECT " + fName1 + "(lst), " + fName2 + "(st), " + fName3 + "(mp) FROM %s WHERE key = 1"),
|
||||
row(list, set, map));
|
||||
|
||||
for (int version = Server.VERSION_2; version <= maxProtocolVersion; version++)
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fName1 + "(lst), " + fName2 + "(st), " + fName3 + "(mp) FROM %s WHERE key = 1"),
|
||||
row(list, set, map));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptTupleType() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, tup frozen<tuple<double, text, int, boolean>>)");
|
||||
|
||||
String fName = createFunction(KEYSPACE_PER_TEST, "tuple<double, text, int, boolean>",
|
||||
"CREATE FUNCTION %s( tup tuple<double, text, int, boolean> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS tuple<double, text, int, boolean> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$tup;$$;");
|
||||
|
||||
Object t = tuple(1d, "foo", 2, true);
|
||||
|
||||
execute("INSERT INTO %s (key, tup) VALUES (1, ?)", t);
|
||||
|
||||
assertRows(execute("SELECT tup FROM %s WHERE key = 1"),
|
||||
row(t));
|
||||
|
||||
assertRows(execute("SELECT " + fName + "(tup) FROM %s WHERE key = 1"),
|
||||
row(t));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptTupleTypeCollection() throws Throwable
|
||||
{
|
||||
String tupleTypeDef = "tuple<double, list<double>, set<text>, map<int, boolean>>";
|
||||
createTable("CREATE TABLE %s (key int primary key, tup frozen<" + tupleTypeDef + ">)");
|
||||
|
||||
String fTup1 = createFunction(KEYSPACE_PER_TEST, tupleTypeDef,
|
||||
"CREATE FUNCTION %s( tup " + tupleTypeDef + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS tuple<double, list<double>, set<text>, map<int, boolean>> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" tup;$$;");
|
||||
String fTup2 = createFunction(KEYSPACE_PER_TEST, tupleTypeDef,
|
||||
"CREATE FUNCTION %s( tup " + tupleTypeDef + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" tup.getDouble(0);$$;");
|
||||
String fTup3 = createFunction(KEYSPACE_PER_TEST, tupleTypeDef,
|
||||
"CREATE FUNCTION %s( tup " + tupleTypeDef + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS list<double> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" tup.getList(1, java.lang.Double.class);$$;");
|
||||
String fTup4 = createFunction(KEYSPACE_PER_TEST, tupleTypeDef,
|
||||
"CREATE FUNCTION %s( tup " + tupleTypeDef + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS set<text> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" tup.getSet(2, java.lang.String.class);$$;");
|
||||
String fTup5 = createFunction(KEYSPACE_PER_TEST, tupleTypeDef,
|
||||
"CREATE FUNCTION %s( tup " + tupleTypeDef + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS map<int, boolean> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" tup.getMap(3, java.lang.Integer.class, java.lang.Boolean.class);$$;");
|
||||
|
||||
List<Double> list = Arrays.asList(1d, 2d, 3d);
|
||||
Set<String> set = new TreeSet<>(Arrays.asList("one", "three", "two"));
|
||||
Map<Integer, Boolean> map = new TreeMap<>();
|
||||
map.put(1, true);
|
||||
map.put(2, false);
|
||||
map.put(3, true);
|
||||
|
||||
Object t = tuple(1d, list, set, map);
|
||||
|
||||
execute("INSERT INTO %s (key, tup) VALUES (1, ?)", t);
|
||||
|
||||
assertRows(execute("SELECT " + fTup1 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(t));
|
||||
assertRows(execute("SELECT " + fTup2 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(1d));
|
||||
assertRows(execute("SELECT " + fTup3 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(list));
|
||||
assertRows(execute("SELECT " + fTup4 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(set));
|
||||
assertRows(execute("SELECT " + fTup5 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(map));
|
||||
|
||||
// same test - but via native protocol
|
||||
TupleType tType = TupleType.of(DataType.cdouble(),
|
||||
DataType.list(DataType.cdouble()),
|
||||
DataType.set(DataType.text()),
|
||||
DataType.map(DataType.cint(), DataType.cboolean()));
|
||||
TupleValue tup = tType.newValue(1d, list, set, map);
|
||||
for (int version = Server.VERSION_2; version <= maxProtocolVersion; version++)
|
||||
{
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fTup1 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(tup));
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fTup2 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(1d));
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fTup3 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(list));
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fTup4 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(set));
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fTup5 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(map));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptUserType() throws Throwable
|
||||
{
|
||||
String type = createType("CREATE TYPE %s (txt text, i int)");
|
||||
|
||||
createTable("CREATE TABLE %s (key int primary key, udt frozen<" + type + ">)");
|
||||
|
||||
String fUdt1 = createFunction(KEYSPACE, type,
|
||||
"CREATE FUNCTION %s( udt " + type + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS " + type + ' ' +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" udt;$$;");
|
||||
String fUdt2 = createFunction(KEYSPACE, type,
|
||||
"CREATE FUNCTION %s( udt " + type + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS text " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" udt.getString(\"txt\");$$;");
|
||||
String fUdt3 = createFunction(KEYSPACE, type,
|
||||
"CREATE FUNCTION %s( udt " + type + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS int " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" udt.getInt(\"i\");$$;");
|
||||
|
||||
execute("INSERT INTO %s (key, udt) VALUES (1, {txt: 'one', i:1})");
|
||||
|
||||
UntypedResultSet rows = execute("SELECT " + fUdt1 + "(udt) FROM %s WHERE key = 1");
|
||||
Assert.assertEquals(1, rows.size());
|
||||
assertRows(execute("SELECT " + fUdt2 + "(udt) FROM %s WHERE key = 1"),
|
||||
row("one"));
|
||||
assertRows(execute("SELECT " + fUdt3 + "(udt) FROM %s WHERE key = 1"),
|
||||
row(1));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptUTCollections() throws Throwable
|
||||
{
|
||||
String type = createType("CREATE TYPE %s (txt text, i int)");
|
||||
|
||||
createTable(String.format("CREATE TABLE %%s " +
|
||||
"(key int primary key, lst list<frozen<%s>>, st set<frozen<%s>>, mp map<int, frozen<%s>>)",
|
||||
type, type, type));
|
||||
|
||||
String fName = createFunction(KEYSPACE, "list<frozen<" + type + ">>",
|
||||
"CREATE FUNCTION %s( lst list<frozen<" + type + ">> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS text " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" lst.get(1).getString(\"txt\");$$;");
|
||||
createFunctionOverload(fName, "set<frozen<" + type + ">>",
|
||||
"CREATE FUNCTION %s( st set<frozen<" + type + ">> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS text " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" st.iterator().next().getString(\"txt\");$$;");
|
||||
createFunctionOverload(fName, "map<int, frozen<" + type + ">>",
|
||||
"CREATE FUNCTION %s( mp map<int, frozen<" + type + ">> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS text " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" mp.get(java.lang.Integer.valueOf(3)).getString(\"txt\");$$;");
|
||||
|
||||
execute("INSERT INTO %s (key, lst, st, mp) values (1, " +
|
||||
// list<frozen<UDT>>
|
||||
"[ {txt: 'one', i:1}, {txt: 'three', i:1}, {txt: 'one', i:1} ] , " +
|
||||
// set<frozen<UDT>>
|
||||
"{ {txt: 'one', i:1}, {txt: 'three', i:3}, {txt: 'two', i:2} }, " +
|
||||
// map<int, frozen<UDT>>
|
||||
"{ 1: {txt: 'one', i:1}, 2: {txt: 'one', i:3}, 3: {txt: 'two', i:2} })");
|
||||
|
||||
assertRows(execute("SELECT " + fName + "(lst) FROM %s WHERE key = 1"),
|
||||
row("three"));
|
||||
assertRows(execute("SELECT " + fName + "(st) FROM %s WHERE key = 1"),
|
||||
row("one"));
|
||||
assertRows(execute("SELECT " + fName + "(mp) FROM %s WHERE key = 1"),
|
||||
row("two"));
|
||||
|
||||
String cqlSelect = "SELECT " + fName + "(lst), " + fName + "(st), " + fName + "(mp) FROM %s WHERE key = 1";
|
||||
assertRows(execute(cqlSelect),
|
||||
row("three", "one", "two"));
|
||||
|
||||
// same test - but via native protocol
|
||||
for (int version = Server.VERSION_2; version <= maxProtocolVersion; version++)
|
||||
assertRowsNet(version,
|
||||
executeNet(version, cqlSelect),
|
||||
row("three", "one", "two"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptFunction() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, val double)");
|
||||
|
||||
String functionBody = '\n' +
|
||||
" Math.sin(val);\n";
|
||||
|
||||
String fName = createFunction(KEYSPACE, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS '" + functionBody + "';");
|
||||
|
||||
FunctionName fNameName = parseFunctionName(fName);
|
||||
|
||||
assertRows(execute("SELECT language, body FROM system_schema.functions WHERE keyspace_name=? AND function_name=?",
|
||||
fNameName.keyspace, fNameName.name),
|
||||
row("javascript", functionBody));
|
||||
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 1, 1d);
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 2, 2d);
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 3, 3d);
|
||||
assertRows(execute("SELECT key, val, " + fName + "(val) FROM %s"),
|
||||
row(1, 1d, Math.sin(1d)),
|
||||
row(2, 2d, Math.sin(2d)),
|
||||
row(3, 3d, Math.sin(3d))
|
||||
);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptBadReturnType() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, val double)");
|
||||
|
||||
String fName = createFunction(KEYSPACE, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS '\"string\";';");
|
||||
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 1, 1d);
|
||||
// throws IRE with ClassCastException
|
||||
assertInvalidMessage("Invalid value for CQL type double", "SELECT key, val, " + fName + "(val) FROM %s");
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptThrow() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, val double)");
|
||||
|
||||
String fName = createFunction(KEYSPACE, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS 'throw \"fool\";';");
|
||||
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 1, 1d);
|
||||
// throws IRE with ScriptException
|
||||
assertInvalidThrowMessage("fool", FunctionExecutionException.class,
|
||||
"SELECT key, val, " + fName + "(val) FROM %s");
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testScriptReturnTypeCasting() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, val double)");
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 1, 1d);
|
||||
|
||||
Object[][] variations = {
|
||||
new Object[] { "true", "boolean", true },
|
||||
new Object[] { "false", "boolean", false },
|
||||
new Object[] { "100", "tinyint", (byte)100 },
|
||||
new Object[] { "100.", "tinyint", (byte)100 },
|
||||
new Object[] { "100", "smallint", (short)100 },
|
||||
new Object[] { "100.", "smallint", (short)100 },
|
||||
new Object[] { "100", "int", 100 },
|
||||
new Object[] { "100.", "int", 100 },
|
||||
new Object[] { "100", "double", 100d },
|
||||
new Object[] { "100.", "double", 100d },
|
||||
new Object[] { "100", "bigint", 100L },
|
||||
new Object[] { "100.", "bigint", 100L },
|
||||
new Object[] { "100", "varint", BigInteger.valueOf(100L) },
|
||||
new Object[] { "100.", "varint", BigInteger.valueOf(100L) },
|
||||
new Object[] { "parseInt(\"100\");", "decimal", BigDecimal.valueOf(100d) },
|
||||
new Object[] { "100.", "decimal", BigDecimal.valueOf(100d) },
|
||||
};
|
||||
|
||||
for (Object[] variation : variations)
|
||||
{
|
||||
Object functionBody = variation[0];
|
||||
Object returnType = variation[1];
|
||||
Object expectedResult = variation[2];
|
||||
|
||||
String fName = createFunction(KEYSPACE, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS " +returnType + ' ' +
|
||||
"LANGUAGE javascript " +
|
||||
"AS '" + functionBody + ";';");
|
||||
assertRows(execute("SELECT key, val, " + fName + "(val) FROM %s"),
|
||||
row(1, 1d, expectedResult));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testScriptParamReturnTypes() throws Throwable
|
||||
{
|
||||
UUID ruuid = UUID.randomUUID();
|
||||
UUID tuuid = UUIDGen.getTimeUUID();
|
||||
|
||||
createTable("CREATE TABLE %s (key int primary key, " +
|
||||
"tival tinyint, sival smallint, ival int, lval bigint, fval float, dval double, vval varint, ddval decimal, " +
|
||||
"timval time, dtval date, tsval timestamp, uval uuid, tuval timeuuid)");
|
||||
execute("INSERT INTO %s (key, tival, sival, ival, lval, fval, dval, vval, ddval, timval, dtval, tsval, uval, tuval) VALUES " +
|
||||
"(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", 1,
|
||||
(byte)1, (short)1, 1, 1L, 1f, 1d, BigInteger.valueOf(1L), BigDecimal.valueOf(1d), 1L, Integer.MAX_VALUE, new Date(1), ruuid, tuuid);
|
||||
|
||||
Object[][] variations = {
|
||||
new Object[] { "tinyint", "tival", (byte)1, (byte)2 },
|
||||
new Object[] { "smallint", "sival", (short)1, (short)2 },
|
||||
new Object[] { "int", "ival", 1, 2 },
|
||||
new Object[] { "bigint", "lval", 1L, 2L },
|
||||
new Object[] { "float", "fval", 1f, 2f },
|
||||
new Object[] { "double", "dval", 1d, 2d },
|
||||
new Object[] { "varint", "vval", BigInteger.valueOf(1L), BigInteger.valueOf(2L) },
|
||||
new Object[] { "decimal", "ddval", BigDecimal.valueOf(1d), BigDecimal.valueOf(2d) },
|
||||
new Object[] { "time", "timval", 1L, 2L },
|
||||
};
|
||||
|
||||
for (Object[] variation : variations)
|
||||
{
|
||||
Object type = variation[0];
|
||||
Object col = variation[1];
|
||||
Object expected1 = variation[2];
|
||||
Object expected2 = variation[3];
|
||||
String fName = createFunction(KEYSPACE, type.toString(),
|
||||
"CREATE OR REPLACE FUNCTION %s(val " + type + ") " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS " + type + ' ' +
|
||||
"LANGUAGE javascript " +
|
||||
"AS 'val+1;';");
|
||||
assertRows(execute("SELECT key, " + col + ", " + fName + '(' + col + ") FROM %s"),
|
||||
row(1, expected1, expected2));
|
||||
}
|
||||
|
||||
variations = new Object[][] {
|
||||
new Object[] { "timestamp","tsval", new Date(1), new Date(1) },
|
||||
new Object[] { "uuid", "uval", ruuid, ruuid },
|
||||
new Object[] { "timeuuid", "tuval", tuuid, tuuid },
|
||||
new Object[] { "date", "dtval", Integer.MAX_VALUE, Integer.MAX_VALUE },
|
||||
};
|
||||
|
||||
for (Object[] variation : variations)
|
||||
{
|
||||
Object type = variation[0];
|
||||
Object col = variation[1];
|
||||
Object expected1 = variation[2];
|
||||
Object expected2 = variation[3];
|
||||
String fName = createFunction(KEYSPACE, type.toString(),
|
||||
"CREATE OR REPLACE FUNCTION %s(val " + type + ") " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS " + type + ' ' +
|
||||
"LANGUAGE javascript " +
|
||||
"AS 'val;';");
|
||||
assertRows(execute("SELECT key, " + col + ", " + fName + '(' + col + ") FROM %s"),
|
||||
row(1, expected1, expected2));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -17,8 +17,6 @@
|
|||
*/
|
||||
package org.apache.cassandra.cql3.validation.entities;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.math.BigInteger;
|
||||
import java.nio.ByteBuffer;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Arrays;
|
||||
|
|
@ -29,26 +27,28 @@ import java.util.Set;
|
|||
import java.util.TreeMap;
|
||||
import java.util.TreeSet;
|
||||
import java.util.UUID;
|
||||
import java.security.AccessControlException;
|
||||
|
||||
import org.junit.Assert;
|
||||
import org.junit.BeforeClass;
|
||||
import org.junit.Test;
|
||||
|
||||
import com.datastax.driver.core.*;
|
||||
import com.datastax.driver.core.exceptions.InvalidQueryException;
|
||||
import org.apache.cassandra.config.DatabaseDescriptor;
|
||||
import org.apache.cassandra.config.Schema;
|
||||
import org.apache.cassandra.cql3.CQLTester;
|
||||
import org.apache.cassandra.cql3.QueryProcessor;
|
||||
import org.apache.cassandra.cql3.UntypedResultSet;
|
||||
import org.apache.cassandra.config.Config;
|
||||
import org.apache.cassandra.cql3.functions.FunctionName;
|
||||
import org.apache.cassandra.cql3.functions.UDFunction;
|
||||
import org.apache.cassandra.cql3.CQLTester;
|
||||
import org.apache.cassandra.cql3.functions.UDHelper;
|
||||
import org.apache.cassandra.db.marshal.CollectionType;
|
||||
import org.apache.cassandra.dht.ByteOrderedPartitioner;
|
||||
import org.apache.cassandra.exceptions.FunctionExecutionException;
|
||||
import org.apache.cassandra.exceptions.InvalidRequestException;
|
||||
import org.apache.cassandra.schema.KeyspaceMetadata;
|
||||
import org.apache.cassandra.service.ClientState;
|
||||
import org.apache.cassandra.service.StorageService;
|
||||
import org.apache.cassandra.service.ClientWarn;
|
||||
import org.apache.cassandra.transport.Event;
|
||||
import org.apache.cassandra.transport.Server;
|
||||
import org.apache.cassandra.transport.messages.ResultMessage;
|
||||
|
|
@ -1521,326 +1521,6 @@ public class UFTest extends CQLTester
|
|||
row("three", "one", "two"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptSimpleCollections() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, lst list<double>, st set<text>, mp map<int, boolean>)");
|
||||
|
||||
String fName1 = createFunction(KEYSPACE_PER_TEST, "list<double>",
|
||||
"CREATE FUNCTION %s( lst list<double> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS list<double> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS 'lst;';");
|
||||
String fName2 = createFunction(KEYSPACE_PER_TEST, "set<text>",
|
||||
"CREATE FUNCTION %s( st set<text> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS set<text> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS 'st;';");
|
||||
String fName3 = createFunction(KEYSPACE_PER_TEST, "map<int, boolean>",
|
||||
"CREATE FUNCTION %s( mp map<int, boolean> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS map<int, boolean> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS 'mp;';");
|
||||
|
||||
List<Double> list = Arrays.asList(1d, 2d, 3d);
|
||||
Set<String> set = new TreeSet<>(Arrays.asList("one", "three", "two"));
|
||||
Map<Integer, Boolean> map = new TreeMap<>();
|
||||
map.put(1, true);
|
||||
map.put(2, false);
|
||||
map.put(3, true);
|
||||
|
||||
execute("INSERT INTO %s (key, lst, st, mp) VALUES (1, ?, ?, ?)", list, set, map);
|
||||
|
||||
assertRows(execute("SELECT lst, st, mp FROM %s WHERE key = 1"),
|
||||
row(list, set, map));
|
||||
|
||||
assertRows(execute("SELECT " + fName1 + "(lst), " + fName2 + "(st), " + fName3 + "(mp) FROM %s WHERE key = 1"),
|
||||
row(list, set, map));
|
||||
|
||||
for (int version = Server.VERSION_2; version <= maxProtocolVersion; version++)
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fName1 + "(lst), " + fName2 + "(st), " + fName3 + "(mp) FROM %s WHERE key = 1"),
|
||||
row(list, set, map));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptTupleType() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, tup frozen<tuple<double, text, int, boolean>>)");
|
||||
|
||||
String fName = createFunction(KEYSPACE_PER_TEST, "tuple<double, text, int, boolean>",
|
||||
"CREATE FUNCTION %s( tup tuple<double, text, int, boolean> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS tuple<double, text, int, boolean> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$tup;$$;");
|
||||
|
||||
Object t = tuple(1d, "foo", 2, true);
|
||||
|
||||
execute("INSERT INTO %s (key, tup) VALUES (1, ?)", t);
|
||||
|
||||
assertRows(execute("SELECT tup FROM %s WHERE key = 1"),
|
||||
row(t));
|
||||
|
||||
assertRows(execute("SELECT " + fName + "(tup) FROM %s WHERE key = 1"),
|
||||
row(t));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptTupleTypeCollection() throws Throwable
|
||||
{
|
||||
String tupleTypeDef = "tuple<double, list<double>, set<text>, map<int, boolean>>";
|
||||
createTable("CREATE TABLE %s (key int primary key, tup frozen<" + tupleTypeDef + ">)");
|
||||
|
||||
String fTup1 = createFunction(KEYSPACE_PER_TEST, tupleTypeDef,
|
||||
"CREATE FUNCTION %s( tup " + tupleTypeDef + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS tuple<double, list<double>, set<text>, map<int, boolean>> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" tup;$$;");
|
||||
String fTup2 = createFunction(KEYSPACE_PER_TEST, tupleTypeDef,
|
||||
"CREATE FUNCTION %s( tup " + tupleTypeDef + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" tup.getDouble(0);$$;");
|
||||
String fTup3 = createFunction(KEYSPACE_PER_TEST, tupleTypeDef,
|
||||
"CREATE FUNCTION %s( tup " + tupleTypeDef + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS list<double> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" tup.getList(1, java.lang.Class.forName(\"java.lang.Double\"));$$;");
|
||||
String fTup4 = createFunction(KEYSPACE_PER_TEST, tupleTypeDef,
|
||||
"CREATE FUNCTION %s( tup " + tupleTypeDef + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS set<text> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" tup.getSet(2, java.lang.Class.forName(\"java.lang.String\"));$$;");
|
||||
String fTup5 = createFunction(KEYSPACE_PER_TEST, tupleTypeDef,
|
||||
"CREATE FUNCTION %s( tup " + tupleTypeDef + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS map<int, boolean> " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" tup.getMap(3, java.lang.Class.forName(\"java.lang.Integer\"), java.lang.Class.forName(\"java.lang.Boolean\"));$$;");
|
||||
|
||||
List<Double> list = Arrays.asList(1d, 2d, 3d);
|
||||
Set<String> set = new TreeSet<>(Arrays.asList("one", "three", "two"));
|
||||
Map<Integer, Boolean> map = new TreeMap<>();
|
||||
map.put(1, true);
|
||||
map.put(2, false);
|
||||
map.put(3, true);
|
||||
|
||||
Object t = tuple(1d, list, set, map);
|
||||
|
||||
execute("INSERT INTO %s (key, tup) VALUES (1, ?)", t);
|
||||
|
||||
assertRows(execute("SELECT " + fTup1 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(t));
|
||||
assertRows(execute("SELECT " + fTup2 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(1d));
|
||||
assertRows(execute("SELECT " + fTup3 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(list));
|
||||
assertRows(execute("SELECT " + fTup4 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(set));
|
||||
assertRows(execute("SELECT " + fTup5 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(map));
|
||||
|
||||
// same test - but via native protocol
|
||||
TupleType tType = TupleType.of(DataType.cdouble(),
|
||||
DataType.list(DataType.cdouble()),
|
||||
DataType.set(DataType.text()),
|
||||
DataType.map(DataType.cint(), DataType.cboolean()));
|
||||
TupleValue tup = tType.newValue(1d, list, set, map);
|
||||
for (int version = Server.VERSION_2; version <= maxProtocolVersion; version++)
|
||||
{
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fTup1 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(tup));
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fTup2 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(1d));
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fTup3 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(list));
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fTup4 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(set));
|
||||
assertRowsNet(version,
|
||||
executeNet(version, "SELECT " + fTup5 + "(tup) FROM %s WHERE key = 1"),
|
||||
row(map));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptUserType() throws Throwable
|
||||
{
|
||||
String type = createType("CREATE TYPE %s (txt text, i int)");
|
||||
|
||||
createTable("CREATE TABLE %s (key int primary key, udt frozen<" + type + ">)");
|
||||
|
||||
String fUdt1 = createFunction(KEYSPACE, type,
|
||||
"CREATE FUNCTION %s( udt " + type + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS " + type + " " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" udt;$$;");
|
||||
String fUdt2 = createFunction(KEYSPACE, type,
|
||||
"CREATE FUNCTION %s( udt " + type + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS text " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" udt.getString(\"txt\");$$;");
|
||||
String fUdt3 = createFunction(KEYSPACE, type,
|
||||
"CREATE FUNCTION %s( udt " + type + " ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS int " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" udt.getInt(\"i\");$$;");
|
||||
|
||||
execute("INSERT INTO %s (key, udt) VALUES (1, {txt: 'one', i:1})");
|
||||
|
||||
UntypedResultSet rows = execute("SELECT " + fUdt1 + "(udt) FROM %s WHERE key = 1");
|
||||
Assert.assertEquals(1, rows.size());
|
||||
assertRows(execute("SELECT " + fUdt2 + "(udt) FROM %s WHERE key = 1"),
|
||||
row("one"));
|
||||
assertRows(execute("SELECT " + fUdt3 + "(udt) FROM %s WHERE key = 1"),
|
||||
row(1));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptUTCollections() throws Throwable
|
||||
{
|
||||
String type = createType("CREATE TYPE %s (txt text, i int)");
|
||||
|
||||
createTable(String.format("CREATE TABLE %%s " +
|
||||
"(key int primary key, lst list<frozen<%s>>, st set<frozen<%s>>, mp map<int, frozen<%s>>)",
|
||||
type, type, type));
|
||||
|
||||
String fName = createFunction(KEYSPACE, "list<frozen<" + type + ">>",
|
||||
"CREATE FUNCTION %s( lst list<frozen<" + type + ">> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS text " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" lst.get(1).getString(\"txt\");$$;");
|
||||
createFunctionOverload(fName, "set<frozen<" + type + ">>",
|
||||
"CREATE FUNCTION %s( st set<frozen<" + type + ">> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS text " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" st.iterator().next().getString(\"txt\");$$;");
|
||||
createFunctionOverload(fName, "map<int, frozen<" + type + ">>",
|
||||
"CREATE FUNCTION %s( mp map<int, frozen<" + type + ">> ) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS text " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS $$" +
|
||||
" mp.get(java.lang.Integer.valueOf(3)).getString(\"txt\");$$;");
|
||||
|
||||
execute("INSERT INTO %s (key, lst, st, mp) values (1, " +
|
||||
// list<frozen<UDT>>
|
||||
"[ {txt: 'one', i:1}, {txt: 'three', i:1}, {txt: 'one', i:1} ] , " +
|
||||
// set<frozen<UDT>>
|
||||
"{ {txt: 'one', i:1}, {txt: 'three', i:3}, {txt: 'two', i:2} }, " +
|
||||
// map<int, frozen<UDT>>
|
||||
"{ 1: {txt: 'one', i:1}, 2: {txt: 'one', i:3}, 3: {txt: 'two', i:2} })");
|
||||
|
||||
assertRows(execute("SELECT " + fName + "(lst) FROM %s WHERE key = 1"),
|
||||
row("three"));
|
||||
assertRows(execute("SELECT " + fName + "(st) FROM %s WHERE key = 1"),
|
||||
row("one"));
|
||||
assertRows(execute("SELECT " + fName + "(mp) FROM %s WHERE key = 1"),
|
||||
row("two"));
|
||||
|
||||
String cqlSelect = "SELECT " + fName + "(lst), " + fName + "(st), " + fName + "(mp) FROM %s WHERE key = 1";
|
||||
assertRows(execute(cqlSelect),
|
||||
row("three", "one", "two"));
|
||||
|
||||
// same test - but via native protocol
|
||||
for (int version = Server.VERSION_2; version <= maxProtocolVersion; version++)
|
||||
assertRowsNet(version,
|
||||
executeNet(version, cqlSelect),
|
||||
row("three", "one", "two"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptFunction() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, val double)");
|
||||
|
||||
String functionBody = '\n' +
|
||||
" Math.sin(val);\n";
|
||||
|
||||
String fName = createFunction(KEYSPACE, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS '" + functionBody + "';");
|
||||
|
||||
FunctionName fNameName = parseFunctionName(fName);
|
||||
|
||||
assertRows(execute("SELECT language, body FROM system_schema.functions WHERE keyspace_name=? AND function_name=?",
|
||||
fNameName.keyspace, fNameName.name),
|
||||
row("javascript", functionBody));
|
||||
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 1, 1d);
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 2, 2d);
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 3, 3d);
|
||||
assertRows(execute("SELECT key, val, " + fName + "(val) FROM %s"),
|
||||
row(1, 1d, Math.sin(1d)),
|
||||
row(2, 2d, Math.sin(2d)),
|
||||
row(3, 3d, Math.sin(3d))
|
||||
);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptBadReturnType() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, val double)");
|
||||
|
||||
String fName = createFunction(KEYSPACE, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS '\"string\";';");
|
||||
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 1, 1d);
|
||||
// throws IRE with ClassCastException
|
||||
assertInvalidMessage("Invalid value for CQL type double", "SELECT key, val, " + fName + "(val) FROM %s");
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testJavascriptThrow() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, val double)");
|
||||
|
||||
String fName = createFunction(KEYSPACE, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS 'throw \"fool\";';");
|
||||
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 1, 1d);
|
||||
// throws IRE with ScriptException
|
||||
assertInvalidThrowMessage("fool", FunctionExecutionException.class,
|
||||
"SELECT key, val, " + fName + "(val) FROM %s");
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testDuplicateArgNames() throws Throwable
|
||||
{
|
||||
|
|
@ -1874,113 +1554,6 @@ public class UFTest extends CQLTester
|
|||
"AS 'question for 42?';");
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testScriptReturnTypeCasting() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, val double)");
|
||||
execute("INSERT INTO %s (key, val) VALUES (?, ?)", 1, 1d);
|
||||
|
||||
Object[][] variations = {
|
||||
new Object[] { "true", "boolean", true },
|
||||
new Object[] { "false", "boolean", false },
|
||||
new Object[] { "100", "tinyint", (byte)100 },
|
||||
new Object[] { "100.", "tinyint", (byte)100 },
|
||||
new Object[] { "100", "smallint", (short)100 },
|
||||
new Object[] { "100.", "smallint", (short)100 },
|
||||
new Object[] { "100", "int", 100 },
|
||||
new Object[] { "100.", "int", 100 },
|
||||
new Object[] { "100", "double", 100d },
|
||||
new Object[] { "100.", "double", 100d },
|
||||
new Object[] { "100", "bigint", 100L },
|
||||
new Object[] { "100.", "bigint", 100L },
|
||||
new Object[] { "100", "varint", BigInteger.valueOf(100L) },
|
||||
new Object[] { "100.", "varint", BigInteger.valueOf(100L) },
|
||||
new Object[] { "parseInt(\"100\");", "decimal", BigDecimal.valueOf(100d) },
|
||||
new Object[] { "100.", "decimal", BigDecimal.valueOf(100d) },
|
||||
};
|
||||
|
||||
for (Object[] variation : variations)
|
||||
{
|
||||
Object functionBody = variation[0];
|
||||
Object returnType = variation[1];
|
||||
Object expectedResult = variation[2];
|
||||
|
||||
String fName = createFunction(KEYSPACE, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS " +returnType + ' ' +
|
||||
"LANGUAGE javascript " +
|
||||
"AS '" + functionBody + ";';");
|
||||
assertRows(execute("SELECT key, val, " + fName + "(val) FROM %s"),
|
||||
row(1, 1d, expectedResult));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testScriptParamReturnTypes() throws Throwable
|
||||
{
|
||||
UUID ruuid = UUID.randomUUID();
|
||||
UUID tuuid = UUIDGen.getTimeUUID();
|
||||
|
||||
createTable("CREATE TABLE %s (key int primary key, " +
|
||||
"tival tinyint, sival smallint, ival int, lval bigint, fval float, dval double, vval varint, ddval decimal, " +
|
||||
"timval time, dtval date, tsval timestamp, uval uuid, tuval timeuuid)");
|
||||
execute("INSERT INTO %s (key, tival, sival, ival, lval, fval, dval, vval, ddval, timval, dtval, tsval, uval, tuval) VALUES " +
|
||||
"(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", 1,
|
||||
(byte)1, (short)1, 1, 1L, 1f, 1d, BigInteger.valueOf(1L), BigDecimal.valueOf(1d), 1L, Integer.MAX_VALUE, new Date(1), ruuid, tuuid);
|
||||
|
||||
Object[][] variations = {
|
||||
new Object[] { "tinyint", "tival", (byte)1, (byte)2 },
|
||||
new Object[] { "smallint", "sival", (short)1, (short)2 },
|
||||
new Object[] { "int", "ival", 1, 2 },
|
||||
new Object[] { "bigint", "lval", 1L, 2L },
|
||||
new Object[] { "float", "fval", 1f, 2f },
|
||||
new Object[] { "double", "dval", 1d, 2d },
|
||||
new Object[] { "varint", "vval", BigInteger.valueOf(1L), BigInteger.valueOf(2L) },
|
||||
new Object[] { "decimal", "ddval", BigDecimal.valueOf(1d), BigDecimal.valueOf(2d) },
|
||||
new Object[] { "time", "timval", 1L, 2L },
|
||||
};
|
||||
|
||||
for (Object[] variation : variations)
|
||||
{
|
||||
Object type = variation[0];
|
||||
Object col = variation[1];
|
||||
Object expected1 = variation[2];
|
||||
Object expected2 = variation[3];
|
||||
String fName = createFunction(KEYSPACE, type.toString(),
|
||||
"CREATE OR REPLACE FUNCTION %s(val " + type + ") " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS " + type + ' ' +
|
||||
"LANGUAGE javascript " +
|
||||
"AS 'val+1;';");
|
||||
assertRows(execute("SELECT key, " + col + ", " + fName + '(' + col + ") FROM %s"),
|
||||
row(1, expected1, expected2));
|
||||
}
|
||||
|
||||
variations = new Object[][] {
|
||||
new Object[] { "timestamp","tsval", new Date(1), new Date(1) },
|
||||
new Object[] { "uuid", "uval", ruuid, ruuid },
|
||||
new Object[] { "timeuuid", "tuval", tuuid, tuuid },
|
||||
new Object[] { "date", "dtval", Integer.MAX_VALUE, Integer.MAX_VALUE },
|
||||
};
|
||||
|
||||
for (Object[] variation : variations)
|
||||
{
|
||||
Object type = variation[0];
|
||||
Object col = variation[1];
|
||||
Object expected1 = variation[2];
|
||||
Object expected2 = variation[3];
|
||||
String fName = createFunction(KEYSPACE, type.toString(),
|
||||
"CREATE OR REPLACE FUNCTION %s(val " + type + ") " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS " + type + ' ' +
|
||||
"LANGUAGE javascript " +
|
||||
"AS 'val;';");
|
||||
assertRows(execute("SELECT key, " + col + ", " + fName + '(' + col + ") FROM %s"),
|
||||
row(1, expected1, expected2));
|
||||
}
|
||||
}
|
||||
|
||||
static class TypesTestDef
|
||||
{
|
||||
final String udfType;
|
||||
|
|
@ -2217,7 +1790,6 @@ public class UFTest extends CQLTester
|
|||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testFunctionWithFrozenSetType() throws Throwable
|
||||
{
|
||||
|
|
@ -2586,9 +2158,220 @@ public class UFTest extends CQLTester
|
|||
assertRows(execute("SELECT " + fNameBRN + "(bval) FROM %s"), row(ByteBufferUtil.EMPTY_BYTE_BUFFER));
|
||||
assertRows(execute("SELECT " + fNameBCC + "(bval) FROM %s"), row(ByteBufferUtil.EMPTY_BYTE_BUFFER));
|
||||
assertRows(execute("SELECT " + fNameBCN + "(bval) FROM %s"), row(ByteBufferUtil.EMPTY_BYTE_BUFFER));
|
||||
assertRows(execute("SELECT " + fNameIRC + "(empty_int) FROM %s"), row(new Object[]{null}));
|
||||
assertRows(execute("SELECT " + fNameIRN + "(empty_int) FROM %s"), row(new Object[]{null}));
|
||||
assertRows(execute("SELECT " + fNameIRC + "(empty_int) FROM %s"), row(new Object[]{ null }));
|
||||
assertRows(execute("SELECT " + fNameIRN + "(empty_int) FROM %s"), row(new Object[]{ null }));
|
||||
assertRows(execute("SELECT " + fNameICC + "(empty_int) FROM %s"), row(0));
|
||||
assertRows(execute("SELECT " + fNameICN + "(empty_int) FROM %s"), row(new Object[]{null}));
|
||||
assertRows(execute("SELECT " + fNameICN + "(empty_int) FROM %s"), row(new Object[]{ null }));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testSecurityPermissions() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, dval double)");
|
||||
execute("INSERT INTO %s (key, dval) VALUES (?, ?)", 1, 1d);
|
||||
|
||||
// Java UDFs
|
||||
|
||||
try
|
||||
{
|
||||
String fName = createFunction(KEYSPACE_PER_TEST, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE JAVA\n" +
|
||||
"AS 'System.getProperty(\"foo.bar.baz\"); return 0d;';");
|
||||
execute("SELECT " + fName + "(dval) FROM %s WHERE key=1");
|
||||
Assert.fail();
|
||||
}
|
||||
catch (FunctionExecutionException e)
|
||||
{
|
||||
assertAccessControlException("System.getProperty(\"foo.bar.baz\"); return 0d;", e);
|
||||
}
|
||||
|
||||
String[][] typesAndSources =
|
||||
{
|
||||
{"", "try { Class.forName(\"" + UDHelper.class.getName() + "\"); } catch (Exception e) { throw new RuntimeException(e); } return 0d;"},
|
||||
{"sun.misc.Unsafe", "sun.misc.Unsafe.getUnsafe(); return 0d;"},
|
||||
{"", "try { Class.forName(\"sun.misc.Unsafe\"); } catch (Exception e) { throw new RuntimeException(e); } return 0d;"},
|
||||
{"java.nio.file.FileSystems", "try {" +
|
||||
" java.nio.file.FileSystems.getDefault(); return 0d;" +
|
||||
"} catch (Exception t) {" +
|
||||
" throw new RuntimeException(t);" +
|
||||
'}'},
|
||||
{"java.nio.channels.FileChannel", "try {" +
|
||||
" java.nio.channels.FileChannel.open(java.nio.file.FileSystems.getDefault().getPath(\"/etc/passwd\")).close(); return 0d;" +
|
||||
"} catch (Exception t) {" +
|
||||
" throw new RuntimeException(t);" +
|
||||
'}'},
|
||||
{"java.nio.channels.SocketChannel", "try {" +
|
||||
" java.nio.channels.SocketChannel.open().close(); return 0d;" +
|
||||
"} catch (Exception t) {" +
|
||||
" throw new RuntimeException(t);" +
|
||||
'}'},
|
||||
{"java.io.FileInputStream", "try {" +
|
||||
" new java.io.FileInputStream(\"./foobar\").close(); return 0d;" +
|
||||
"} catch (Exception t) {" +
|
||||
" throw new RuntimeException(t);" +
|
||||
'}'},
|
||||
{"java.lang.Runtime", "try {" +
|
||||
" java.lang.Runtime.getRuntime(); return 0d;" +
|
||||
"} catch (Exception t) {" +
|
||||
" throw new RuntimeException(t);" +
|
||||
'}'},
|
||||
{"org.apache.cassandra.service.StorageService",
|
||||
"try {" +
|
||||
" org.apache.cassandra.service.StorageService v = org.apache.cassandra.service.StorageService.instance; v.isInShutdownHook(); return 0d;" +
|
||||
"} catch (Exception t) {" +
|
||||
" throw new RuntimeException(t);" +
|
||||
'}'},
|
||||
{"java.net.ServerSocket", "try {" +
|
||||
" new java.net.ServerSocket().bind(); return 0d;" +
|
||||
"} catch (Exception t) {" +
|
||||
" throw new RuntimeException(t);" +
|
||||
'}'},
|
||||
{"java.io.FileOutputStream","try {" +
|
||||
" new java.io.FileOutputStream(\".foo\"); return 0d;" +
|
||||
"} catch (Exception t) {" +
|
||||
" throw new RuntimeException(t);" +
|
||||
'}'},
|
||||
{"java.lang.Runtime", "try {" +
|
||||
" java.lang.Runtime.getRuntime().exec(\"/tmp/foo\"); return 0d;" +
|
||||
"} catch (Exception t) {" +
|
||||
" throw new RuntimeException(t);" +
|
||||
'}'}
|
||||
};
|
||||
|
||||
for (String[] typeAndSource : typesAndSources)
|
||||
{
|
||||
assertInvalidMessage(typeAndSource[0] + " cannot be resolved",
|
||||
"CREATE OR REPLACE FUNCTION " + KEYSPACE + ".invalid_class_access(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE JAVA\n" +
|
||||
"AS '" + typeAndSource[1] + "';");
|
||||
}
|
||||
|
||||
// JavaScript UDFs
|
||||
|
||||
try
|
||||
{
|
||||
String fName = createFunction(KEYSPACE_PER_TEST, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS 'org.apache.cassandra.service.StorageService.instance.isInShutdownHook(); 0;';");
|
||||
execute("SELECT " + fName + "(dval) FROM %s WHERE key=1");
|
||||
Assert.fail("Javascript security check failed");
|
||||
}
|
||||
catch (FunctionExecutionException e)
|
||||
{
|
||||
assertAccessControlException("", e);
|
||||
}
|
||||
|
||||
String[] javascript =
|
||||
{
|
||||
"java.lang.management.ManagmentFactory.getThreadMXBean(); 0;",
|
||||
"new java.io.FileInputStream(\"/tmp/foo\"); 0;",
|
||||
"new java.io.FileOutputStream(\"/tmp/foo\"); 0;",
|
||||
"java.nio.file.FileSystems.getDefault().createFileExclusively(\"./foo_bar_baz\"); 0;",
|
||||
"java.nio.channels.FileChannel.open(java.nio.file.FileSystems.getDefault().getPath(\"/etc/passwd\")); 0;",
|
||||
"java.nio.channels.SocketChannel.open(); 0;",
|
||||
"new java.net.ServerSocket().bind(null); 0;",
|
||||
"var thread = new java.lang.Thread(); thread.start(); 0;",
|
||||
"java.lang.System.getProperty(\"foo.bar.baz\"); 0;",
|
||||
"java.lang.Class.forName(\"java.lang.System\"); 0;",
|
||||
"java.lang.Runtime.getRuntime().exec(\"/tmp/foo\"); 0;",
|
||||
"java.lang.Runtime.getRuntime().loadLibrary(\"foobar\"); 0;",
|
||||
"java.lang.Runtime.getRuntime().loadLibrary(\"foobar\"); 0;",
|
||||
// TODO these (ugly) calls are still possible - these can consume CPU (as one could do with an evil loop, too)
|
||||
// "java.lang.Runtime.getRuntime().traceMethodCalls(true); 0;",
|
||||
// "java.lang.Runtime.getRuntime().gc(); 0;",
|
||||
// "java.lang.Runtime.getRuntime(); 0;",
|
||||
};
|
||||
|
||||
for (String script : javascript)
|
||||
{
|
||||
try
|
||||
{
|
||||
String fName = createFunction(KEYSPACE_PER_TEST, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE javascript\n" +
|
||||
"AS '" + script + "';");
|
||||
execute("SELECT " + fName + "(dval) FROM %s WHERE key=1");
|
||||
Assert.fail("Javascript security check failed: " + script);
|
||||
}
|
||||
catch (FunctionExecutionException e)
|
||||
{
|
||||
assertAccessControlException(script, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static void assertAccessControlException(String script, FunctionExecutionException e)
|
||||
{
|
||||
for (Throwable t = e; t != null && t != t.getCause(); t = t.getCause())
|
||||
if (t instanceof AccessControlException)
|
||||
return;
|
||||
Assert.fail("no AccessControlException for " + script + " (got " + e + ')');
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testAmokUDF() throws Throwable
|
||||
{
|
||||
createTable("CREATE TABLE %s (key int primary key, dval double)");
|
||||
execute("INSERT INTO %s (key, dval) VALUES (?, ?)", 1, 1d);
|
||||
|
||||
long udfWarnTimeout = DatabaseDescriptor.getUserDefinedFunctionWarnTimeout();
|
||||
long udfFailTimeout = DatabaseDescriptor.getUserDefinedFunctionFailTimeout();
|
||||
try
|
||||
{
|
||||
// short timeout
|
||||
DatabaseDescriptor.setUserDefinedFunctionWarnTimeout(1);
|
||||
DatabaseDescriptor.setUserDefinedFunctionFailTimeout(100);
|
||||
// don't kill the unit test... - default policy is "die"
|
||||
DatabaseDescriptor.setUserFunctionTimeoutPolicy(Config.UserFunctionTimeoutPolicy.ignore);
|
||||
|
||||
ClientWarn.captureWarnings();
|
||||
String fName = createFunction(KEYSPACE_PER_TEST, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE JAVA\n" +
|
||||
"AS 'long t=System.currentTimeMillis()+20; while (t>System.currentTimeMillis()) { }; return 0d;'");
|
||||
execute("SELECT " + fName + "(dval) FROM %s WHERE key=1");
|
||||
List<String> warnings = ClientWarn.getWarnings();
|
||||
Assert.assertNotNull(warnings);
|
||||
Assert.assertFalse(warnings.isEmpty());
|
||||
ClientWarn.resetWarnings();
|
||||
|
||||
// Java UDF
|
||||
|
||||
fName = createFunction(KEYSPACE_PER_TEST, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE JAVA\n" +
|
||||
"AS 'long t=System.currentTimeMillis()+300; while (t>System.currentTimeMillis()) { }; return 0d;';");
|
||||
assertInvalidMessage("ran longer than 100ms", "SELECT " + fName + "(dval) FROM %s WHERE key=1");
|
||||
|
||||
// Javascript UDF
|
||||
|
||||
fName = createFunction(KEYSPACE_PER_TEST, "double",
|
||||
"CREATE OR REPLACE FUNCTION %s(val double) " +
|
||||
"RETURNS NULL ON NULL INPUT " +
|
||||
"RETURNS double " +
|
||||
"LANGUAGE JAVASCRIPT\n" +
|
||||
"AS 'var t=java.lang.System.currentTimeMillis()+300; while (t>java.lang.System.currentTimeMillis()) { }; 0;';");
|
||||
assertInvalidMessage("ran longer than 100ms", "SELECT " + fName + "(dval) FROM %s WHERE key=1");
|
||||
}
|
||||
finally
|
||||
{
|
||||
// reset to defaults
|
||||
DatabaseDescriptor.setUserDefinedFunctionWarnTimeout(udfWarnTimeout);
|
||||
DatabaseDescriptor.setUserDefinedFunctionFailTimeout(udfFailTimeout);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in New Issue