import argparse from urllib.parse import urlparse import boto3 from pyspark.sql import SparkSession def validate(spark, urls): head_url, *tail_urls = urls head_df = spark.read.parquet(head_url).cache() for (url, df) in list(map(lambda u: (u, spark.read.parquet(u)), tail_urls)): # Compare schemas assert head_df.schema == df.schema, \ 'Schemas are not equal: {first_url} <=> {second_url}'.format( first_url=head_url, second_url=url) # Compare dataframes assert head_df.cache().subtract(df).union(df.subtract(head_df)).count() == 0, \ 'Data are not equal: {first_url} <=> {second_url}'.format( first_url=head_url, second_url=url) def list_s3_keys(bucket, prefix): keys = [] s3 = boto3.client('s3') kwargs = { 'Bucket': bucket, 'Prefix': prefix } while True: resp = s3.list_objects_v2(**kwargs) for obj in resp['Contents']: keys.append(obj['Key']) try: kwargs['ContinuationToken'] = resp['NextContinuationToken'] except KeyError: break return keys def explode_urls(urls): exploded = [] for url in set(urls): if url.endswith('*'): url_parts = urlparse(url.replace('*', '')) bucket = url_parts.hostname prefix = url_parts.path[1:] prefix_size = len(prefix[-1:].split('/')) folders = map(lambda x: x.split('/', prefix_size + 1)[prefix_size] + '/', list_s3_keys(bucket, prefix)) exploded.extend(map(lambda x: url.replace('*', x), set(folders))) else: exploded.append(url) return exploded def main(cli_args): urls = explode_urls(cli_args.urls) assert len(urls) > 1, 'At least 2 urls expected. But only {x} provided.'.format(x=len(urls)) spark = SparkSession.builder \ .appName('Delphi Spark POC Validator') \ .config('spark.jars.packages', 'org.apache.hadoop:hadoop-aws:2.7.3') \ .config('spark.executor.extraJavaOptions', '-XX:+UseG1GC') \ .config('spark.hadoop.mapreduce.fileoutputcommitter.algorithm.version', '2') \ .config('spark.hadoop.mapreduce.fileoutputcommitter.cleanup-failures.ignored', 'true') \ .config('spark.hadoop.parquet.enable.summary-metadata', 'false') \ .config('spark.sql.parquet.mergeSchema', 'false') \ .config('spark.sql.hive.metastorePartitionPruning', 'true') \ .getOrCreate() validate(spark, urls) if __name__ == '__main__': parser = argparse.ArgumentParser() parser.add_argument('--urls', '-u', dest='urls', nargs='+', type=str, required=True) main(parser.parse_args())