progress on mapping data, finding clusters, probably inefficient
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user