progress on mapping data, finding clusters, probably inefficient

This commit is contained in:
nitowa
2022-08-23 15:12:52 -04:00
parent 883c81b786
commit 5b9ec0da6a
27 changed files with 431 additions and 50 deletions
+76
View File
@@ -0,0 +1,76 @@
from gc import collect
from sqlite3 import Row
from typing import Iterable
from operator import add
from pyspark.sql import SparkSession
from pyspark.sql import functions as F
spark = SparkSession.builder \
.appName('SparkCassandraApp') \
.config('spark.cassandra.connection.host', 'localhost') \
.config('spark.cassandra.connection.port', '9042') \
.config('spark.cassandra.output.consistency.level', 'ONE') \
.config("spark.sql.extensions", "com.datastax.spark.connector.CassandraSparkExtensions") \
.config('directJoinSetting', 'on') \
.master('spark://osboxes:7077') \
.getOrCreate()
spark.conf.set("spark.sql.catalog.myCatalog",
"com.datastax.spark.connector.datasource.CassandraCatalog")
tx_addr_groups = spark.read.table("myCatalog.distributedunionfind.transactions") \
.groupBy("tx_id") \
.agg(F.collect_set('address').alias('addresses')) \
.toLocalIterator()
def insertCluster (row):
addrs: Iterable[str] = row['addresses']
df = spark.createDataFrame(map(lambda addr: (addr, addrs[0]), addrs), schema=['address', 'parent'])
df.writeTo("myCatalog.distributedunionfind.clusters").overwrite()
"""
tuple structure:
Row => Row(parent=addr, addresses=list[addr]
Iterable[str] => list[addr]
"""
def find(data: tuple[Row, Iterable[str]]):
cluster = data[0]
tx = data[1]
clusteraddresses = cluster['addresses'] + [cluster['parent']]
if any(x in tx for x in clusteraddresses):
return cluster['parent']
else:
return None
for addr_group in tx_addr_groups:
clusters_df = spark.read.table("myCatalog.distributedunionfind.clusters")
clusters = clusters_df \
.groupBy("parent") \
.agg(F.collect_set('address').alias('addresses'))
if (clusters.count() == 0):
insertCluster(addr_group)
continue
df = clusters.rdd \
.map(lambda cluster: (cluster, addr_group['addresses'])) \
.map(find) \
.filter(lambda x: x != None) \
.collect()
if(len(df) == 0):
insertCluster(addr_group)
continue
print(addr_group)
print(df)
break