from pyspark import SparkConf from pyspark import SparkContext from pyspark import sql import config as conf_values class Tester: def __init__(self): self.setup_spark() result = {} def setup_spark(self): conf = SparkConf() \ .setAppName(conf_values.spark_app) \ .setMaster(conf_values.spark_master) self.sc = SparkContext(conf=conf) self.sc._jsc.hadoopConfiguration().set("fs.s3n.awsAccessKeyId", conf_values.aws_access_key) self.sc._jsc.hadoopConfiguration().set("fs.s3n.awsSecretAccessKey", conf_values.aws_secret_key) self.sql_context = sql.SQLContext(self.sc) def load_table(self, table_name, query=''): table = getattr(conf_values, table_name) print ('SOURCE ', table['source']) query_str = '' if query == '': query_str = table['query'] else: query_str = query return self.sql_context.read \ .format("com.databricks.spark.redshift") \ .option("url", table['source']) \ .option("query", query_str) \ .option("tempdir", "s3n://dev-cucumbers/YouTubeMonthly/spark/") \ .load() def test_me(selfi, info): print ('TEST ME', info) def test_comparison(self, info): if len(info['columns']) == 0: print ('INFO ', (info['tables'])[0]) df_1 = self.load_table(info['tables'][0]) df_2 = self.load_table(info['tables'][1]) df_1.show() df_2.show() result_df = df_1.subtract(df_2) result_df.printSchema() result_df.show() else: columns = ','.join([col for col in info['columns']]) print 'columns ', columns df_1 = self.load_table(info['tables'][0], 'SELECT {} from {}'.format(columns, info['tables'][0])) df_2 = self.load_table(info['tables'][0], 'SELECT {} from {}'.format(columns, info['tables'][1])) result_df = df_1.subtract(df_2) result_df.printSchema() result_df.show() def test_aggregate(self, info): for table in info['tables']: # assume column is numeric for column in info['columns']: df = self.load_table(table, 'SELECT {} from {}'.format(column, table)) df.describe().show() def get_tests(self): for test_name in conf_values.tests: if test_name == 'test_comparison': #self.test_me(getattr(conf_values, test_name)) self.test_comparison(getattr(conf_values, test_name)) if test_name == 'test_aggregate': self.test_aggregate(getattr(conf_values, test_name)) if __name__ == "__main__": t = Tester() t.get_tests()