Skip to content
Closed
Show file tree
Hide file tree
Changes from 6 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
2 changes: 1 addition & 1 deletion core/src/main/scala/org/apache/spark/SparkConf.scala
Original file line number Diff line number Diff line change
Expand Up @@ -248,7 +248,7 @@ class SparkConf(loadDefaults: Boolean) extends Cloneable with Logging {
* - This will throw an exception is the config is not optional and the value is not set.
*/
private[spark] def get[T](entry: ConfigEntry[T]): T = {
entry.readFrom(this)
entry.readFrom(settings, getenv)
}

/**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,11 +17,33 @@

package org.apache.spark.internal.config

import java.util.{Map => JMap}

import scala.util.matching.Regex

import org.apache.spark.SparkConf

/**
* An entry contains all meta information for a configuration.
*
* Config options created using this feature support variable expansion. If the config value
* contains variable references of the form "${prefix:variableName}", the reference will be replaced
* with the value of the variable depending on the prefix. The prefix can be one of:
*
* - no prefix: if the config key starts with "spark", looks for the value in the Spark config
* - system: looks for the value in the system properties
* - env: looks for the value in the environment
*
* So referencing "${spark.master}" will look for the value of "spark.master" in the Spark
* configuration, while referencing "${env:MASTER}" will read the value from the "MASTER"
* environment variable.
*
* For known Spark configuration keys (i.e. those created using `ConfigBuilder`), references
* will also consider the default value when it exists.
*
* If the reference cannot be resolved, the original string will be retained. Variable expansion
* only applies to user-provided values, not to default values.
*
* @param key the key for the configuration
* @param defaultValue the default value for the configuration
* @param valueConverter how to convert a string to the value. It should throw an exception if the
Expand All @@ -42,17 +64,27 @@ private[spark] abstract class ConfigEntry[T] (
val doc: String,
val isPublic: Boolean) {

import ConfigEntry._

registerEntry(this)

def defaultValueString: String

def readFrom(conf: SparkConf): T
def readFrom(conf: JMap[String, String], getenv: String => String): T

// This is used by SQLConf, since it doesn't use SparkConf to store settings and thus cannot
// use readFrom().
def defaultValue: Option[T] = None

override def toString: String = {
s"ConfigEntry(key=$key, defaultValue=$defaultValueString, doc=$doc, public=$isPublic)"
}

protected def readAndExpand(
conf: JMap[String, String],
getenv: String => String,
usedRefs: Set[String] = Set()): Option[String] = {
Option(conf.get(key)).map(expand(_, conf, getenv, usedRefs))
}

}

private class ConfigEntryWithDefault[T] (
Expand All @@ -68,8 +100,8 @@ private class ConfigEntryWithDefault[T] (

override def defaultValueString: String = stringConverter(_defaultValue)

override def readFrom(conf: SparkConf): T = {
conf.getOption(key).map(valueConverter).getOrElse(_defaultValue)
def readFrom(conf: JMap[String, String], getenv: String => String): T = {
readAndExpand(conf, getenv).map(valueConverter).getOrElse(_defaultValue)
}

}
Expand All @@ -88,7 +120,9 @@ private[spark] class OptionalConfigEntry[T](

override def defaultValueString: String = "<undefined>"

override def readFrom(conf: SparkConf): Option[T] = conf.getOption(key).map(rawValueConverter)
override def readFrom(conf: JMap[String, String], getenv: String => String): Option[T] = {
readAndExpand(conf, getenv).map(rawValueConverter)
}

}

Expand All @@ -99,13 +133,66 @@ private class FallbackConfigEntry[T] (
key: String,
doc: String,
isPublic: Boolean,
private val fallback: ConfigEntry[T])
private[config] val fallback: ConfigEntry[T])
extends ConfigEntry[T](key, fallback.valueConverter, fallback.stringConverter, doc, isPublic) {

override def defaultValueString: String = s"<value of ${fallback.key}>"

override def readFrom(conf: SparkConf): T = {
conf.getOption(key).map(valueConverter).getOrElse(fallback.readFrom(conf))
override def readFrom(conf: JMap[String, String], getenv: String => String): T = {
Option(conf.get(key)).map(valueConverter).getOrElse(fallback.readFrom(conf, getenv))
}

}

private object ConfigEntry {

private val knownConfigs = new java.util.concurrent.ConcurrentHashMap[String, ConfigEntry[_]]()

private val REF_RE = "\\$\\{(?:(\\w+?):)?(\\S+?)\\}".r

def registerEntry(entry: ConfigEntry[_]): Unit = {
val existing = knownConfigs.putIfAbsent(entry.key, entry)
require(existing == null, s"Config entry ${entry.key} already registered!")
}

def findEntry(key: String): ConfigEntry[_] = knownConfigs.get(key)

/**
* Expand the `value` according to the rules explained in ConfigEntry.
*/
def expand(
value: String,
conf: JMap[String, String],
getenv: String => String,
usedRefs: Set[String]): String = {
REF_RE.replaceAllIn(value, { m =>
val prefix = m.group(1)
val name = m.group(2)
val replacement = prefix match {
case null =>
require(!usedRefs.contains(name), s"Circular reference in $value: $name")
if (name.startsWith("spark.")) {
Option(findEntry(name))
.flatMap(_.readAndExpand(conf, getenv, usedRefs = usedRefs + name))
.orElse(Option(conf.get(name)))
.orElse(defaultValueString(name))
} else {
None
}
case "system" => sys.props.get(name)
case "env" => Option(getenv(name))
case _ => throw new IllegalArgumentException(s"Invalid prefix: $prefix")

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should this throw?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yeah, I have to take a look at this. Throwing might actually break some existing code in SparkHadoopUtil.

}
Regex.quoteReplacement(replacement.getOrElse(m.matched))
})
}

private def defaultValueString(key: String): Option[String] = {
findEntry(key) match {
case e: ConfigEntryWithDefault[_] => Some(e.defaultValueString)
case e: FallbackConfigEntry[_] => defaultValueString(e.fallback.key)
case _ => None
}
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -19,53 +19,60 @@ package org.apache.spark.internal.config

import java.util.concurrent.TimeUnit

import scala.collection.JavaConverters._
import scala.collection.mutable.HashMap

import org.apache.spark.{SparkConf, SparkFunSuite}
import org.apache.spark.network.util.ByteUnit

class ConfigEntrySuite extends SparkFunSuite {

private val PREFIX = "spark.ConfigEntrySuite."

private def testKey(name: String): String = s"$PREFIX.$name"

test("conf entry: int") {
val conf = new SparkConf()
val iConf = ConfigBuilder("spark.int").intConf.createWithDefault(1)
val iConf = ConfigBuilder(testKey("int")).intConf.createWithDefault(1)
assert(conf.get(iConf) === 1)
conf.set(iConf, 2)
assert(conf.get(iConf) === 2)
}

test("conf entry: long") {
val conf = new SparkConf()
val lConf = ConfigBuilder("spark.long").longConf.createWithDefault(0L)
val lConf = ConfigBuilder(testKey("long")).longConf.createWithDefault(0L)
conf.set(lConf, 1234L)
assert(conf.get(lConf) === 1234L)
}

test("conf entry: double") {
val conf = new SparkConf()
val dConf = ConfigBuilder("spark.double").doubleConf.createWithDefault(0.0)
val dConf = ConfigBuilder(testKey("double")).doubleConf.createWithDefault(0.0)
conf.set(dConf, 20.0)
assert(conf.get(dConf) === 20.0)
}

test("conf entry: boolean") {
val conf = new SparkConf()
val bConf = ConfigBuilder("spark.boolean").booleanConf.createWithDefault(false)
val bConf = ConfigBuilder(testKey("boolean")).booleanConf.createWithDefault(false)
assert(!conf.get(bConf))
conf.set(bConf, true)
assert(conf.get(bConf))
}

test("conf entry: optional") {
val conf = new SparkConf()
val optionalConf = ConfigBuilder("spark.optional").intConf.createOptional
val optionalConf = ConfigBuilder(testKey("optional")).intConf.createOptional
assert(conf.get(optionalConf) === None)
conf.set(optionalConf, 1)
assert(conf.get(optionalConf) === Some(1))
}

test("conf entry: fallback") {
val conf = new SparkConf()
val parentConf = ConfigBuilder("spark.int").intConf.createWithDefault(1)
val confWithFallback = ConfigBuilder("spark.fallback").fallbackConf(parentConf)
val parentConf = ConfigBuilder(testKey("parent")).intConf.createWithDefault(1)
val confWithFallback = ConfigBuilder(testKey("fallback")).fallbackConf(parentConf)
assert(conf.get(confWithFallback) === 1)
conf.set(confWithFallback, 2)
assert(conf.get(parentConf) === 1)
Expand All @@ -74,23 +81,25 @@ class ConfigEntrySuite extends SparkFunSuite {

test("conf entry: time") {
val conf = new SparkConf()
val time = ConfigBuilder("spark.time").timeConf(TimeUnit.SECONDS).createWithDefaultString("1h")
val time = ConfigBuilder(testKey("time")).timeConf(TimeUnit.SECONDS)
.createWithDefaultString("1h")
assert(conf.get(time) === 3600L)
conf.set(time.key, "1m")
assert(conf.get(time) === 60L)
}

test("conf entry: bytes") {
val conf = new SparkConf()
val bytes = ConfigBuilder("spark.bytes").bytesConf(ByteUnit.KiB).createWithDefaultString("1m")
val bytes = ConfigBuilder(testKey("bytes")).bytesConf(ByteUnit.KiB)
.createWithDefaultString("1m")
assert(conf.get(bytes) === 1024L)
conf.set(bytes.key, "1k")
assert(conf.get(bytes) === 1L)
}

test("conf entry: string seq") {
val conf = new SparkConf()
val seq = ConfigBuilder("spark.seq").stringConf.toSequence.createWithDefault(Seq())
val seq = ConfigBuilder(testKey("seq")).stringConf.toSequence.createWithDefault(Seq())
conf.set(seq.key, "1,,2, 3 , , 4")
assert(conf.get(seq) === Seq("1", "2", "3", "4"))
conf.set(seq, Seq("1", "2"))
Expand All @@ -99,7 +108,7 @@ class ConfigEntrySuite extends SparkFunSuite {

test("conf entry: int seq") {
val conf = new SparkConf()
val seq = ConfigBuilder("spark.seq").intConf.toSequence.createWithDefault(Seq())
val seq = ConfigBuilder(testKey("intSeq")).intConf.toSequence.createWithDefault(Seq())
conf.set(seq.key, "1,,2, 3 , , 4")
assert(conf.get(seq) === Seq(1, 2, 3, 4))
conf.set(seq, Seq(1, 2))
Expand All @@ -108,7 +117,7 @@ class ConfigEntrySuite extends SparkFunSuite {

test("conf entry: transformation") {
val conf = new SparkConf()
val transformationConf = ConfigBuilder("spark.transformation")
val transformationConf = ConfigBuilder(testKey("transformation"))
.stringConf
.transform(_.toLowerCase())
.createWithDefault("FOO")
Expand All @@ -120,7 +129,7 @@ class ConfigEntrySuite extends SparkFunSuite {

test("conf entry: valid values check") {
val conf = new SparkConf()
val enum = ConfigBuilder("spark.enum")
val enum = ConfigBuilder(testKey("enum"))
.stringConf
.checkValues(Set("a", "b", "c"))
.createWithDefault("a")
Expand All @@ -138,7 +147,7 @@ class ConfigEntrySuite extends SparkFunSuite {

test("conf entry: conversion error") {
val conf = new SparkConf()
val conversionTest = ConfigBuilder("spark.conversionTest").doubleConf.createOptional
val conversionTest = ConfigBuilder(testKey("conversionTest")).doubleConf.createOptional
conf.set(conversionTest.key, "abc")
val conversionError = intercept[IllegalArgumentException] {
conf.get(conversionTest)
Expand All @@ -148,8 +157,72 @@ class ConfigEntrySuite extends SparkFunSuite {

test("default value handling is null-safe") {
val conf = new SparkConf()
val stringConf = ConfigBuilder("spark.string").stringConf.createWithDefault(null)
val stringConf = ConfigBuilder(testKey("string")).stringConf.createWithDefault(null)
assert(conf.get(stringConf) === null)
}

test("variable expansion") {
val env = Map("ENV1" -> "env1")
val conf = HashMap("spark.value1" -> "value1", "spark.value2" -> "value2")

def getenv(key: String): String = env.getOrElse(key, null)

def expand(value: String): String = ConfigEntry.expand(value, conf.asJava, getenv, Set())

assert(expand("${spark.value1}") === "value1")
assert(expand("spark.value1 is: ${spark.value1}") === "spark.value1 is: value1")
assert(expand("${spark.value1} ${spark.value2}") === "value1 value2")
assert(expand("${spark.value3}") === "${spark.value3}")

// Make sure anything that is not in the "spark." namespace is ignored.
conf("notspark.key") = "value"
assert(expand("${notspark.key}") === "${notspark.key}")

assert(expand("${env:ENV1}") === "env1")
assert(expand("${system:user.name}") === sys.props("user.name"))

val stringConf = ConfigBuilder(testKey("stringForExpansion"))
.stringConf
.createWithDefault("string1")
val optionalConf = ConfigBuilder(testKey("optionForExpansion"))
.stringConf
.createOptional
val intConf = ConfigBuilder(testKey("intForExpansion"))
.intConf
.createWithDefault(42)
val fallbackConf = ConfigBuilder(testKey("fallbackForExpansion"))
.fallbackConf(intConf)

assert(expand("${" + stringConf.key + "}") === "string1")
assert(expand("${" + optionalConf.key + "}") === "${" + optionalConf.key + "}")
assert(expand("${" + intConf.key + "}") === "42")
assert(expand("${" + fallbackConf.key + "}") === "42")

conf(optionalConf.key) = "string2"
assert(expand("${" + optionalConf.key + "}") === "string2")

conf(fallbackConf.key) = "84"
assert(expand("${" + fallbackConf.key + "}") === "84")

assert(expand("${spark.value1") === "${spark.value1")

// Chained references.
val conf1 = ConfigBuilder(testKey("conf1"))
.stringConf
.createWithDefault("value1")
val conf2 = ConfigBuilder(testKey("conf2"))
.stringConf
.createWithDefault("value2")

conf(conf2.key) = "${" + conf1.key + "}"
assert(expand("${" + conf2.key + "}") === conf1.defaultValueString)

// Circular references.
conf(conf1.key) = "${" + conf2.key + "}"
val e = intercept[IllegalArgumentException] {
expand("${" + conf2.key + "}")
}
assert(e.getMessage().contains("Circular"))
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -738,8 +738,7 @@ private[sql] class SQLConf extends Serializable with CatalystConf with Logging {
*/
def getConf[T](entry: ConfigEntry[T]): T = {
require(sqlConfEntries.get(entry.key) == entry, s"$entry is not registered")
Option(settings.get(entry.key)).map(entry.valueConverter).orElse(entry.defaultValue).
getOrElse(throw new NoSuchElementException(entry.key))
entry.readFrom(settings, System.getenv)
}

/**
Expand All @@ -748,7 +747,7 @@ private[sql] class SQLConf extends Serializable with CatalystConf with Logging {
*/
def getConf[T](entry: OptionalConfigEntry[T]): Option[T] = {
require(sqlConfEntries.get(entry.key) == entry, s"$entry is not registered")
Option(settings.get(entry.key)).map(entry.rawValueConverter)
entry.readFrom(settings, System.getenv)
}

/**
Expand Down