Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -496,10 +496,20 @@ trait HasCognitiveServiceInput extends HasURL with HasSubscriptionKey with HasAA

protected val aadHeaderName = "Authorization"

// Header ServiceParams may become sequences during automatic batching. Payload ServiceParams
// continue to use getValueOpt so document-aligned values such as text and language stay batched.
private def getHeaderStringValueOpt(row: Row, param: ServiceParam[String]): Option[String] =
ServiceHeaderValues.stringValue(getValueAnyOpt(row, param), param.name)

private def getHeaderMapValueOpt(
row: Row,
param: ServiceParam[Map[String, String]]): Option[Map[String, String]] =
ServiceHeaderValues.mapValue(getValueAnyOpt(row, param), param.name)

protected def contentType: Row => String = { _ => "application/json" }

protected def getCustomAuthHeader(row: Row): Option[String] = {
getValueOpt(row, CustomAuthHeader)
getHeaderStringValueOpt(row, CustomAuthHeader)
}

// The automatic Fabric fallback is eligible only when the request carries no explicit subscription
Expand All @@ -512,7 +522,7 @@ trait HasCognitiveServiceInput extends HasURL with HasSubscriptionKey with HasAA
// fetches) is never reached when a non-blank embedded api-key/Authorization is present.
private[ml] def lacksExplicitAuthCredential(row: Row): Boolean =
!Seq(subscriptionKey, AADToken, CustomAuthHeader)
.exists(param => getValueOpt(row, param).exists(ServiceAuthHeaders.nonBlank))
.exists(param => getHeaderStringValueOpt(row, param).isDefined)

// The automatic Fabric fallback is the lowest-priority credential. It is supplied by-name to
// ServiceAuthHeaders.build and therefore invoked only when build's precedence chain finds no
Expand All @@ -530,7 +540,7 @@ trait HasCognitiveServiceInput extends HasURL with HasSubscriptionKey with HasAA
}

protected def getCustomHeaders(row: Row): Option[Map[String, String]] = {
getValueOpt(row, customHeaders)
getHeaderMapValueOpt(row, customHeaders)
}

protected def supportsImplicitFabricAuthRetry: Boolean = false
Expand Down Expand Up @@ -571,14 +581,14 @@ trait HasCognitiveServiceInput extends HasURL with HasSubscriptionKey with HasAA
addContentType: Boolean,
fabricFallbackAuthHeader: => Option[String]): ServiceAuthHeaders.Resolution = {
ServiceAuthHeaders.resolve(
getValueOpt(row, subscriptionKey),
getHeaderStringValueOpt(row, subscriptionKey),
subscriptionKeyHeaderName,
aadHeaderName,
getValueOpt(row, AADToken),
getHeaderStringValueOpt(row, AADToken),
getCustomAuthHeader(row),
getCustomHeaders(row),
fabricFallbackAuthHeader,
getValueOpt(row, telemHeaders),
getHeaderMapValueOpt(row, telemHeaders),
if (addContentType) Option(contentType(row)) else None)
}

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
// Copyright (C) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License. See LICENSE in project root for information.

package com.microsoft.azure.synapse.ml.services

private[ml] object ServiceHeaderValues {

private def values(value: Option[Any]): Iterator[Any] = value.iterator.flatMap {
case batch: scala.collection.Seq[_] => batch.iterator
case scalar => Iterator.single(scalar)
}.flatMap(value => Option(value))

def stringValue(value: Option[Any], paramName: String): Option[String] = {
values(value).map {
case stringValue: String => stringValue
case _ => throw new IllegalArgumentException(
s"Header service parameter '$paramName' must reference a String or array<string> column")
}.find(ServiceAuthHeaders.nonBlank)
}

def mapValue(value: Option[Any], paramName: String): Option[Map[String, String]] = {
values(value).map {
case mapValue: scala.collection.Map[_, _] =>
if (mapValue.exists { case (name, value) =>
Option(name).exists(headerName => !headerName.isInstanceOf[String]) ||
Option(value).exists(headerValue => !headerValue.isInstanceOf[String])
}) {
throw invalidMapType(paramName)
}
ServiceAuthHeaders.sanitizeHeaderMap(mapValue.iterator.collect {
case (name: String, headerValue: String) => name -> headerValue
}.toMap)
case _ => throw invalidMapType(paramName)
}.find(_.nonEmpty)
}

private def invalidMapType(paramName: String): IllegalArgumentException = {
new IllegalArgumentException(
s"Header service parameter '$paramName' must reference a map<string,string> " +
"or array<map<string,string>> column")
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,16 @@ package com.microsoft.azure.synapse.ml.services

import com.microsoft.azure.synapse.ml.core.test.base.TestBase
import com.microsoft.azure.synapse.ml.param.ServiceParam
import com.microsoft.azure.synapse.ml.stages.{FixedMiniBatchTransformer, FlattenBatch}
import org.apache.http.entity.AbstractHttpEntity
import org.apache.spark.ml.param.{ParamMap, Params}
import org.apache.spark.sql.Row
import org.apache.spark.sql.catalyst.expressions.GenericRowWithSchema
import org.apache.spark.sql.types.{ArrayType, MapType, StringType, StructType}
import spray.json.DefaultJsonProtocol._

import scala.collection.mutable.ArrayBuffer

private class ServiceParamHarness(override val uid: String = "serviceParamHarness")
extends Params with HasServiceParams {

Expand All @@ -19,6 +24,9 @@ private class ServiceParamHarness(override val uid: String = "serviceParamHarnes
val optionalText: ServiceParam[String] =
new ServiceParam[String](this, "optionalText", "optional text")

val batchedText: ServiceParam[Seq[String]] =
new ServiceParam[Seq[String]](this, "batchedText", "batched text")

val urlVersion: ServiceParam[String] =
new ServiceParam[String](this, "urlVersion", "url version", isURLParam = true)

Expand All @@ -34,6 +42,8 @@ private class ServiceParamHarness(override val uid: String = "serviceParamHarnes

def valueAnyOpt(row: Row, p: ServiceParam[_]): Option[Any] = getValueAnyOpt(row, p)

def valueOpt[T](row: Row, p: ServiceParam[T]): Option[T] = getValueOpt(row, p)

override def copy(extra: ParamMap): Params = this
}

Expand Down Expand Up @@ -186,4 +196,175 @@ class CognitiveServiceBaseSuite extends TestBase {
assert(customHeaders("X-Test") == "1")
assert(customHeaders.contains("x-ai-telemetry-properties"))
}

test("cognitive input helper methods resolve automatically batched string headers") {
val keyInput = new CognitiveInputHarness()
keyInput.setSubscriptionKeyCol("keys")
val keyRow = Seq(Seq(
Option.empty[String], Some(""), Some("sub-key"), Some("other-key")
)).toDF("keys").head()
assert(keyInput.headers(keyRow)("Ocp-Apim-Subscription-Key") == "sub-key")

val aadInput = new CognitiveInputHarness()
aadInput.setAADTokenCol("tokens")
val aadRow = Seq(Seq(
Option.empty[String], Some(""), Some("aad-token"), Some("other-token")
)).toDF("tokens").head()
assert(aadInput.headers(aadRow)("Authorization") == "Bearer aad-token")

val customAuthInput = new CognitiveInputHarness()
customAuthInput.setCustomAuthHeaderCol("authHeaders")
val customAuthRow = Seq(Seq(
Option.empty[String], Some(""), Some("Shared custom-auth"), Some("Shared other-auth")
))
.toDF("authHeaders")
.head()
assert(customAuthInput.headers(customAuthRow)("Authorization") == "Shared custom-auth")
}

test("subscription key columns work through public batching and flattening") {
val input = Seq(
("first", "first-key"),
("second", "second-key")
).toDF("text", "key").coalesce(1)
val batched = new FixedMiniBatchTransformer().setBatchSize(10).transform(input)
val inputBuilder = new CognitiveInputHarness().setSubscriptionKeyCol("key")

assert(batched.head().getAs[scala.collection.Seq[String]]("key") == Seq("first-key", "second-key"))
assert(inputBuilder.headers(batched.head())("Ocp-Apim-Subscription-Key") == "first-key")

val restored = new FlattenBatch().transform(batched)
.select("text", "key")
.as[(String, String)]
.collect()
.toSeq
assert(restored == Seq("first" -> "first-key", "second" -> "second-key"))
}

test("cognitive input helper methods resolve automatically batched map headers") {
val input = new CognitiveInputHarness()
input.setVectorParam(input.customHeaders, "customHeadersCol")
input.setVectorParam(input.telemHeaders, "telemHeadersCol")

val row = Seq((
Seq(Map.empty[String, String], Map("X-Test" -> "1")),
Seq(Map.empty[String, String], Map("X-Telemetry" -> "2"))
)).toDF("customHeadersCol", "telemHeadersCol").head()

val headers = input.headers(row)
assert(headers("X-Test") == "1")
assert(headers("X-Telemetry") == "2")
}

test("header columns accept mutable Spark array representations") {
val input = new CognitiveInputHarness().setSubscriptionKeyCol("keys")
input.setVectorParam(input.customHeaders, "custom")
val schema = new StructType()
.add("keys", ArrayType(StringType))
.add("custom", ArrayType(MapType(StringType, StringType)))
val row = new GenericRowWithSchema(Array[Any](
ArrayBuffer("", "batch-key"),
ArrayBuffer(Map.empty[String, String], Map("X-Test" -> "kept"))
), schema)

val headers = input.headers(row)
assert(headers("Ocp-Apim-Subscription-Key") == "batch-key")
assert(headers("X-Test") == "kept")
}

test("empty automatically batched credentials do not fail header resolution") {
val input = new CognitiveInputHarness()
input.setSubscriptionKeyCol("keys")
input.setAADTokenCol("tokens")
input.setCustomAuthHeaderCol("authHeaders")

val row = Seq((
Seq(Option.empty[String], Some("")),
Seq(Option.empty[String], Some("")),
Seq(Option.empty[String], Some(""))
)).toDF("keys", "tokens", "authHeaders").head()

val headers = input.headers(row)
assert(!headers.contains("Ocp-Apim-Subscription-Key"))
assert(!headers.contains("Authorization"))
}

test("batched credentials preserve auth precedence and lazy Fabric fallback") {
val input = new CognitiveInputHarness().setSubscriptionKeyCol("keys").setAADTokenCol("tokens")
input.setVectorParam(input.customHeaders, "custom")
input.setVectorParam(input.telemHeaders, "telemetry")
val row = Seq((
Seq("", "batch-key"),
Seq("batch-token"),
Seq(Map("Authorization" -> "embedded-token", "X-Custom" -> "kept")),
Seq(Map("AUTHORIZATION" -> "ignored", "Ocp-Apim-Subscription-Key" -> "ignored", "X-Telemetry" -> "kept"))
)).toDF("keys", "tokens", "custom", "telemetry").head()

def unexpectedFallback: Option[String] = throw new AssertionError("Fallback must remain lazy")

val headers = input.buildServiceAuthHeaders(row, addContentType = true,
fabricFallbackAuthHeader = unexpectedFallback)
assert(headers("Ocp-Apim-Subscription-Key") == "batch-key")
assert(!headers.keys.exists(_.equalsIgnoreCase("Authorization")))
assert(headers("X-Custom") == "kept")
assert(headers("X-Telemetry") == "kept")
assert(!input.lacksExplicitAuthCredential(row))
}

test("all-blank batched credentials allow Fabric fallback") {
val input = new CognitiveInputHarness().setSubscriptionKeyCol("keys")
val row = Seq(Seq(Option.empty[String], Some(" "), Some(""))).toDF("keys").head()
assert(input.lacksExplicitAuthCredential(row))
val headers = input.buildServiceAuthHeaders(row, addContentType = false,
fabricFallbackAuthHeader = Some("fallback-token"))
assert(headers("Authorization") == "fallback-token")
assert(!headers.contains("Ocp-Apim-Subscription-Key"))
}

test("batched maps skip null and sanitized-empty maps without merging later maps") {
val input = new CognitiveInputHarness()
input.setVectorParam(input.customHeaders, "custom")
val row = Seq(Seq(
Option.empty[Map[String, String]],
Some(Map("X-Null" -> Option.empty[String].orNull)),
Some(Map("X-First" -> "kept")),
Some(Map("X-Later" -> "ignored"))
)).toDF("custom").head()
val headers = input.headers(row)
assert(headers("X-First") == "kept")
assert(!headers.contains("X-Null"))
assert(!headers.contains("X-Later"))
}

test("invalid automatically batched credential element types fail clearly") {
val input = new CognitiveInputHarness().setSubscriptionKeyCol("keys")
val row = Seq(Seq(1, 2)).toDF("keys").head()

val error = intercept[IllegalArgumentException](input.headers(row))
assert(error.getMessage.contains("subscriptionKey"))
assert(error.getMessage.contains("String or array<string>"))

val mapInput = new CognitiveInputHarness()
mapInput.setVectorParam(mapInput.customHeaders, "customHeadersCol")
val invalidMapRow = Seq(Seq("not-a-map")).toDF("customHeadersCol").head()

val mapError = intercept[IllegalArgumentException](mapInput.headers(invalidMapRow))
assert(mapError.getMessage.contains("customHeaders"))
assert(mapError.getMessage.contains("map<string,string> or array<map<string,string>>"))

val invalidEntryRow = Seq(Seq(Map("sensitive-test-value" -> 123))).toDF("customHeadersCol").head()
val entryError = intercept[IllegalArgumentException](mapInput.headers(invalidEntryRow))
assert(entryError.getMessage.contains("customHeaders"))
assert(!entryError.getMessage.contains("sensitive-test-value"))
assert(!entryError.getMessage.contains("123"))
}

test("batch-aware header resolution does not change payload service parameters") {
val harness = new ServiceParamHarness()
harness.setVectorParam(harness.batchedText, "textCol")
val values = Seq("first", "second")
val row = Seq(values).toDF("textCol").head()

assert(harness.valueOpt(row, harness.batchedText).contains(values))
}
}
Loading
Loading