Skip to content
Open
Original file line number Diff line number Diff line change
Expand Up @@ -364,7 +364,7 @@ class SentimentAnalysisLROSuite extends TransformerFuzzing[AnalyzeTextLongRunnin
override def testObjects(): Seq[TestObject[AnalyzeTextLongRunningOperations]] =
Seq(new TestObject[AnalyzeTextLongRunningOperations](model, df))

override def reader: MLReadable[_] = AnalyzeText
override def reader: MLReadable[_] = AnalyzeTextLongRunningOperations
}


Expand Down Expand Up @@ -602,7 +602,7 @@ class EntityRecognitionLROSuite extends TransformerFuzzing[AnalyzeTextLongRunnin
override def testObjects(): Seq[TestObject[AnalyzeTextLongRunningOperations]] =
Seq(new TestObject[AnalyzeTextLongRunningOperations](model, df))

override def reader: MLReadable[_] = AnalyzeText
override def reader: MLReadable[_] = AnalyzeTextLongRunningOperations
}

class CustomEntityRecognitionSuite extends TransformerFuzzing[AnalyzeTextLongRunningOperations]
Expand Down Expand Up @@ -648,7 +648,7 @@ class CustomEntityRecognitionSuite extends TransformerFuzzing[AnalyzeTextLongRun
override def testObjects(): Seq[TestObject[AnalyzeTextLongRunningOperations]] =
Seq(new TestObject[AnalyzeTextLongRunningOperations](model, df))

override def reader: MLReadable[_] = AnalyzeText
override def reader: MLReadable[_] = AnalyzeTextLongRunningOperations
}


Expand Down Expand Up @@ -697,8 +697,7 @@ class MultiLableClassificationSuite extends TransformerFuzzing[AnalyzeTextLongRu
override def testObjects(): Seq[TestObject[AnalyzeTextLongRunningOperations]] =
Seq(new TestObject[AnalyzeTextLongRunningOperations](model, df))

override def reader: MLReadable[_] = AnalyzeText
override def reader: MLReadable[_] = AnalyzeTextLongRunningOperations
}



7 changes: 7 additions & 0 deletions core/src/main/python/synapse/ml/core/schema/Utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
from pyspark.ml.wrapper import JavaParams
from pyspark.ml.common import inherit_doc, _java2py
from pyspark import SparkContext
from pyspark.sql import SparkSession
from synapse.ml.core.serialize._safe_import import secure_import_class


Expand Down Expand Up @@ -58,6 +59,11 @@ def read(cls):

@inherit_doc
class ComplexParamsMixin(MLReadable):
@classmethod
def read(cls):
"""Returns a reader bound to the active Spark session."""
return JavaMMLReader(cls)

def _transfer_params_from_java(self):
"""
Transforms the embedded com.microsoft.azure.synapse.ml.core.serialize.params from the companion Java object.
Expand Down Expand Up @@ -131,6 +137,7 @@ class JavaMMLReader(JavaMLReader):

def __init__(self, clazz):
super(JavaMMLReader, self).__init__(clazz)
self.session(SparkSession.builder.getOrCreate())

@classmethod
def _java_loader_class(cls, clazz):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

package com.microsoft.azure.synapse.ml.core.serialize

import com.microsoft.azure.synapse.ml.core.utils.DeserializationClassFilter
import com.microsoft.azure.synapse.ml.param.WrappableParam
import org.apache.hadoop.fs.Path
import org.apache.spark.ml.Serializer
Expand All @@ -16,12 +17,36 @@ abstract class ComplexParam[T: TypeTag](parent: Params, name: String, doc: Strin

def ttag: TypeTag[T] = typeTag[T]

/** Class policy for legacy Java object streams. No policy means loading is disabled unless the
* Spark session explicitly opts into trusted legacy deserialization.
*/
protected def deserializationClassFilter: Option[DeserializationClassFilter] = None

protected def supportsUntrustedDeserialization: Boolean = true

def isSafeForUntrustedDeserialization: Boolean = {
supportsUntrustedDeserialization &&
(!Serializer.usesObjectSerializer(ttag.tpe) || deserializationClassFilter.isDefined)
}

def save(obj: T, sparkSession: SparkSession, path: Path, overwrite: Boolean): Unit = {
Serializer.typeToSerializer[T](ttag.tpe, sparkSession).write(obj, path, overwrite)
Serializer.typeToSerializer[T](ttag.tpe, sparkSession, deserializationClassFilter)
.write(obj, path, overwrite)
}

def load(sparkSession: SparkSession, path: Path): T = {
Serializer.typeToSerializer[T](ttag.tpe, sparkSession).read(path)
if (
!isSafeForUntrustedDeserialization &&
!Serializer.trustedLoadEnabled(sparkSession)
) {
throw new SecurityException(
s"Complex parameter $name requires a trusted artifact. Set " +
s"${Serializer.LegacyObjectDeserializationConfig}=true and load through " +
"read.session(sparkSession).load(path), or wrap a native Pipeline load in " +
"Serializer.withTrustedArtifactLoad(sparkSession), only when loading trusted data."
)
}
Serializer.typeToSerializer[T](ttag.tpe, sparkSession, deserializationClassFilter).read(path)
}

override def jsonEncode(value: T): String = {
Expand Down
Loading
Loading