diff --git a/common/workflow-operator/src/main/scala/org/apache/texera/amber/operator/sklearn/SklearnPredictionOpDesc.scala b/common/workflow-operator/src/main/scala/org/apache/texera/amber/operator/sklearn/SklearnPredictionOpDesc.scala index c1d022974ce..206bed80792 100644 --- a/common/workflow-operator/src/main/scala/org/apache/texera/amber/operator/sklearn/SklearnPredictionOpDesc.scala +++ b/common/workflow-operator/src/main/scala/org/apache/texera/amber/operator/sklearn/SklearnPredictionOpDesc.scala @@ -92,8 +92,15 @@ class SklearnPredictionOpDesc extends PythonOperatorDescriptor with StandaloneCo var resultType = AttributeType.STRING val inputSchema = inputSchemas(operatorInfo.inputPorts(1).id) if (groundTruthAttribute != "") { - resultType = - inputSchema.attributes.find(attr => attr.getName == groundTruthAttribute).get.getType + // Exact case: Schema lookups ignore case, but the generated drop() does not. + resultType = inputSchema.attributes + .find(attr => attr.getName == groundTruthAttribute) + .getOrElse( + throw new RuntimeException( + s"Ground Truth column '$groundTruthAttribute' is not in the input table" + ) + ) + .getType } Map( operatorInfo.outputPorts.head.id -> inputSchema diff --git a/common/workflow-operator/src/test/scala/org/apache/texera/amber/operator/sklearn/SklearnPredictionOpDescSpec.scala b/common/workflow-operator/src/test/scala/org/apache/texera/amber/operator/sklearn/SklearnPredictionOpDescSpec.scala index b640e4888df..a376901c436 100644 --- a/common/workflow-operator/src/test/scala/org/apache/texera/amber/operator/sklearn/SklearnPredictionOpDescSpec.scala +++ b/common/workflow-operator/src/test/scala/org/apache/texera/amber/operator/sklearn/SklearnPredictionOpDescSpec.scala @@ -71,14 +71,28 @@ class SklearnPredictionOpDescSpec extends AnyFlatSpec with Matchers { .getType shouldBe AttributeType.INTEGER } - it should "throw when the configured ground-truth attribute is absent from the input schema" in { + it should "name the configured ground-truth attribute when it is absent from the input schema" in { val d = new SklearnPredictionOpDesc d.resultAttribute = "prediction" d.groundTruthAttribute = "missing" val data = Schema().add("feature", AttributeType.STRING) - intercept[NoSuchElementException] { + val e = intercept[RuntimeException] { d.getOutputSchemas(Map(PortIdentity(1) -> data)) } + e.getMessage shouldBe "Ground Truth column 'missing' is not in the input table" + } + + it should "reject a ground-truth attribute that matches an input column only by case" in { + val d = new SklearnPredictionOpDesc + d.resultAttribute = "prediction" + d.groundTruthAttribute = "Label" + val data = Schema() + .add("feature", AttributeType.STRING) + .add("label", AttributeType.INTEGER) + val e = intercept[RuntimeException] { + d.getOutputSchemas(Map(PortIdentity(1) -> data)) + } + e.getMessage shouldBe "Ground Truth column 'Label' is not in the input table" } "SklearnPredictionOpDesc.generatePythonCode" should "emit the model-applying tuple operator" in {