diff --git a/src/main/scala/com/springml/spark/salesforce/DataWriter.scala b/src/main/scala/com/springml/spark/salesforce/DataWriter.scala index 2565dad..ec7813d 100644 --- a/src/main/scala/com/springml/spark/salesforce/DataWriter.scala +++ b/src/main/scala/com/springml/spark/salesforce/DataWriter.scala @@ -28,9 +28,10 @@ import org.apache.spark.sql.SaveMode * It uses Partner External Metadata SOAP API to write the dataset */ class DataWriter ( - val userName: String, + val userName: Option[String], val password: String, - val login: String, + val authToken: Option[String], + val serverUrl: String, val version: String, val datasetName: String, val appName: String @@ -40,7 +41,7 @@ class DataWriter ( def writeMetadata(metaDataJson: String, mode: SaveMode, upsert: Boolean): Option[String] = { - val partnerConnection = createConnection(userName, password, login, version) + val partnerConnection = createConnection(userName, password, authToken, serverUrl, version) val oper = operation(mode, upsert) val sobj = new SObject() @@ -87,7 +88,7 @@ class DataWriter ( sobj.setField("InsightsExternalDataId", metadataId) sobj.setField("PartNumber", partNumber) - val partnerConnection = Utils.createConnection(userName, password, login, version) + val partnerConnection = Utils.createConnection(userName, password, authToken, serverUrl, version) val results = partnerConnection.create(Array(sobj)) val resultSuccess = results.map(saveResult => { @@ -108,7 +109,7 @@ class DataWriter ( def commit(id: String): Boolean = { - val partnerConnection = Utils.createConnection(userName, password, login, version) + val partnerConnection = Utils.createConnection(userName, password, authToken, serverUrl, version) val sobj = new SObject() sobj.setType("InsightsExternalData") diff --git a/src/main/scala/com/springml/spark/salesforce/DefaultSource.scala b/src/main/scala/com/springml/spark/salesforce/DefaultSource.scala index 9986e19..05b55ac 100644 --- a/src/main/scala/com/springml/spark/salesforce/DefaultSource.scala +++ b/src/main/scala/com/springml/spark/salesforce/DefaultSource.scala @@ -17,7 +17,7 @@ package com.springml.spark.salesforce import java.text.SimpleDateFormat -import com.springml.salesforce.wave.api.APIFactory +import com.springml.salesforce.wave.api.{APIFactory, BulkAPI, ForceAPI, WaveAPI} import org.apache.log4j.Logger import org.apache.spark.sql.sources.{BaseRelation, CreatableRelationProvider, RelationProvider, SchemaRelationProvider} import org.apache.spark.sql.types.StructType @@ -50,9 +50,19 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr * */ override def createRelation(sqlContext: SQLContext, parameters: Map[String, String], schema: StructType) = { - val username = param(parameters, "SF_USERNAME", "username") - val password = param(parameters, "SF_PASSWORD", "password") - val login = parameters.getOrElse("login", "https://login.salesforce.com") + val username = optionalParam(parameters, "SF_USERNAME", "username") + val authToken = optionalParam(parameters, "SF_AUTHTOKEN", "authToken") + validateMutualExclusive(username, authToken, "username", "authToken") + var password = "" + var serverUrl = "" + if (username.isDefined) { + password = param(parameters, "SF_PASSWORD", "password") + serverUrl = parameters.getOrElse("login", "https://login.salesforce.com") + } else { + serverUrl = parameters.getOrElse("instanceUrl", "https://login.salesforce.com") + } + + val version = parameters.getOrElse("version", "36.0") val saql = parameters.get("saql") val soql = parameters.get("soql") @@ -71,7 +81,15 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr val inferSchemaFlag = flag(inferSchema, "inferSchema") if (saql.isDefined) { - val waveAPI = APIFactory.getInstance.waveAPI(username, password, login, version) + + var waveAPI:WaveAPI = null + + if (username.isDefined) { + waveAPI = APIFactory.getInstance.waveAPI(username.get, password, serverUrl, version) + } else { + waveAPI = APIFactory.getInstance.waveAPIwAuthToken(authToken.get, serverUrl, version) + } + DatasetRelation(waveAPI, null, saql.get, schema, sqlContext, resultVariable, pageSize.toInt, sampleSize.toInt, encodeFields, inferSchemaFlag, replaceDatasetNameWithId.toBoolean, sdf(dateFormat)) @@ -79,9 +97,14 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr if (replaceDatasetNameWithId.toBoolean) { logger.warn("Ignoring 'replaceDatasetNameWithId' option as it is not applicable to soql") } - - val forceAPI = APIFactory.getInstance.forceAPI(username, password, login, + var forceAPI: ForceAPI = null + if (username.isDefined) { + forceAPI = APIFactory.getInstance.forceAPI(username.get, password, serverUrl, version, Integer.getInteger(pageSize), Integer.getInteger(maxRetry)) + } else { + forceAPI = APIFactory.getInstance.forceAPIwAuthToken(authToken.get, serverUrl, + version, Integer.getInteger(pageSize), Integer.getInteger(maxRetry)) + } DatasetRelation(null, forceAPI, soql.get, schema, sqlContext, null, 0, sampleSize.toInt, encodeFields, inferSchemaFlag, replaceDatasetNameWithId.toBoolean, sdf(dateFormat)) @@ -91,12 +114,23 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr override def createRelation(sqlContext: SQLContext, mode: SaveMode, parameters: Map[String, String], data: DataFrame): BaseRelation = { - val username = param(parameters, "SF_USERNAME", "username") - val password = param(parameters, "SF_PASSWORD", "password") + val username = optionalParam(parameters, "SF_USERNAME", "username") + val authToken = optionalParam(parameters, "SF_AUTHTOKEN", "authToken") + + validateMutualExclusive(username, authToken, "username", "authToken") + + var password = "" + var serverUrl = "" + if (username.isDefined) { + password = param(parameters, "SF_PASSWORD", "password") + serverUrl = parameters.getOrElse("login", "https://login.salesforce.com") + } else { + serverUrl = parameters.getOrElse("instanceUrl", "https://login.salesforce.com") + } + val datasetName = parameters.get("datasetName") val sfObject = parameters.get("sfObject") val appName = parameters.getOrElse("appName", null) - val login = parameters.getOrElse("login", "https://login.salesforce.com") val version = parameters.getOrElse("version", "36.0") val usersMetadataConfig = parameters.get("metadataConfig") val upsert = parameters.getOrElse("upsert", "false") @@ -115,21 +149,22 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr } logger.info("Writing dataframe into Salesforce Wave") - writeInSalesforceWave(username, password, login, version, + writeInSalesforceWave(username, password, authToken, serverUrl, version, datasetName.get, appName, usersMetadataConfig, mode, flag(upsert, "upsert"), flag(monitorJob, "monitorJob"), data, metadataFile) } else { logger.info("Updating Salesforce Object") - updateSalesforceObject(username, password, login, version, sfObject.get, mode, data) + updateSalesforceObject(username, password, authToken, serverUrl, version, sfObject.get, mode, data) } return createReturnRelation(data) } private def updateSalesforceObject( - username: String, + username: Option[String], password: String, - login: String, + authToken: Option[String], + serverUrl: String, version: String, sfObject: String, mode: SaveMode, @@ -141,8 +176,13 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr val repartitionedRDD = Utils.repartition(data.rdd) logger.info("no of partitions after repartitioning is " + repartitionedRDD.partitions.length) - val bulkAPI = APIFactory.getInstance.bulkAPI(username, password, login, version) - val writer = new SFObjectWriter(username, password, login, version, sfObject, mode, csvHeader) + var bulkAPI:BulkAPI = null + if (username.isDefined) { + bulkAPI = APIFactory.getInstance.bulkAPI(username.get, password, serverUrl, version) + } else { + bulkAPI = APIFactory.getInstance.bulkAPIwAuthToken(authToken.get, serverUrl, version) + } + val writer = new SFObjectWriter(username, password, authToken, serverUrl, version, sfObject, mode, csvHeader) logger.info("Writing data") val successfulWrite = writer.writeData(repartitionedRDD) logger.info(s"Writing data was successful was $successfulWrite") @@ -153,9 +193,10 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr } private def writeInSalesforceWave( - username: String, + username: Option[String], password: String, - login: String, + authToken: Option[String], + serverUrl: String, version: String, datasetName: String, appName: String, @@ -165,7 +206,7 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr monitorJob: Boolean, data: DataFrame, metadata: Option[String]) { - val dataWriter = new DataWriter(username, password, login, version, datasetName, appName) + val dataWriter = new DataWriter(username, password, authToken, serverUrl, version, datasetName, appName) val metaDataJson = Utils.metadata(metadata, usersMetadataConfig, data.schema, datasetName) @@ -202,7 +243,7 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr if (monitorJob) { logger.info("Monitoring Job status in Salesforce wave") - if (Utils.monitorJob(writtenId.get, username, password, login, version)) { + if (Utils.monitorJob(writtenId.get, username, password, authToken, serverUrl, version)) { logger.info(s"Successfully dataset $datasetName has been processed in Salesforce Wave") } else { sys.error(s"Upload Job for dataset $datasetName failed in Salesforce Wave. Check Monitor Job in Salesforce Wave for more details") @@ -223,6 +264,17 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr } } + + private def optionalParam(parameters: Map[String, String], envName: String, paramName: String) : Option[String] = { + val envProp = sys.env.get(envName) + if (envProp != null && envProp.isDefined) { + return Option(envProp.get) + } + + val param = parameters.getOrElse(paramName, null) + return Option(param) + } + private def param(parameters: Map[String, String], envName: String, paramName: String) : String = { val envProp = sys.env.get(envName) if (envProp != null && envProp.isDefined) { diff --git a/src/main/scala/com/springml/spark/salesforce/SFObjectWriter.scala b/src/main/scala/com/springml/spark/salesforce/SFObjectWriter.scala index f595dc4..ff5f944 100644 --- a/src/main/scala/com/springml/spark/salesforce/SFObjectWriter.scala +++ b/src/main/scala/com/springml/spark/salesforce/SFObjectWriter.scala @@ -13,9 +13,10 @@ import com.springml.salesforce.wave.util.WaveAPIConstants * Next subsequent columns are fields to be updated */ class SFObjectWriter ( - val username: String, + val username: Option[String], val password: String, - val login: String, + val authToken: Option[String], + val serverUrl: String, val apiVersion: String, val sfObject: String, val mode: SaveMode, @@ -63,7 +64,11 @@ class SFObjectWriter ( } def bulkAPI() : BulkAPI = { - APIFactory.getInstance.bulkAPI(username, password, login, apiVersion) + if (username.isDefined) { + return APIFactory.getInstance.bulkAPI(username.get, password, serverUrl, apiVersion) + } else { + return APIFactory.getInstance.bulkAPIwAuthToken(authToken.get, serverUrl, apiVersion) + } } private def operation(mode: SaveMode): String = { diff --git a/src/main/scala/com/springml/spark/salesforce/Utils.scala b/src/main/scala/com/springml/spark/salesforce/Utils.scala index bcd3cef..133972a 100644 --- a/src/main/scala/com/springml/spark/salesforce/Utils.scala +++ b/src/main/scala/com/springml/spark/salesforce/Utils.scala @@ -37,12 +37,17 @@ import scala.util.Try object Utils extends Serializable { @transient val logger = Logger.getLogger("Utils") - def createConnection(username: String, password: String, - login: String, version: String):PartnerConnection = { + def createConnection(username: Option[String], password: String, authToken: Option[String], + serverUrl: String, version: String):PartnerConnection = { val config = new ConnectorConfig() - config.setUsername(username) - config.setPassword(password) - val endpoint = if (login.endsWith("/")) (login + "services/Soap/u/" + version) else (login + "/services/Soap/u/" + version) + if (username.isDefined) { + config.setUsername(username.get) + config.setPassword(password) + } else { + config.setManualLogin(true) + config.setSessionId(authToken.get) + } + val endpoint = if (serverUrl.endsWith("/")) (serverUrl + "services/Soap/u/" + version) else (serverUrl + "/services/Soap/u/" + version) config.setAuthEndpoint(endpoint) config.setServiceEndpoint(endpoint) Connector.newConnection(config) @@ -150,9 +155,9 @@ object Utils extends Serializable { } } - def monitorJob(objId: String, username: String, password: - String, login: String, version: String) : Boolean = { - var partnerConnection = Utils.createConnection(username, password, login, version) + def monitorJob(objId: String, username: Option[String], password: + String, authToken:Option[String], login: String, version: String) : Boolean = { + var partnerConnection = Utils.createConnection(username, password, authToken, login, version) try { monitorJob(partnerConnection, objId, 500) } catch { @@ -161,7 +166,7 @@ object Utils extends Serializable { logger.info("Error Message from Salesforce Wave " + exMsg) if (exMsg contains "Invalid Session") { logger.info("Session expired. Monitoring Job using new connection") - return monitorJob(objId, username, password, login, version) + return monitorJob(objId, username, password, authToken, login, version) } else { throw uefault }