
public class SuperUnionFind {

	int[] id;
	// Represents the number of pointer changes
	private int findPointers;
	private int unionPointers;
	private int finds;
	private int unions;
	private int pc;
	private int w;
	private int worstCaseUnion;
	private int worstCaseFind;
	private int[] size;
	private int[] rank;
	public SuperUnionFind(int N, int pathCompression, int weighted)
	{
		findPointers = 0;
		unionPointers = 0;
		finds = 0;
		unions = 0;
		pc = pathCompression;
		w = weighted;
		id = new int[N];
		rank = new int[N];
		size = new int[N];
		for (int i = 0; i < id.length; i++)
		{
			id[i] = i;
			rank[i] = 0;
			size[i] = 1;
		}
	}

	public void union(int x, int y)
	{
		unions++;
		int oldPointers = unionPointers;
		
		int f1 = findU(x);
		int f2 = findU(y);
	
		if (w == 1)
		{
			if (size[f1] < size[f2])
			{
				id[f1] = f2;
				size[f2] += size[f1];
			}
			else
			{
				id[f2] = f1;
				size[f1] = f2;		
			}
			unionPointers++;
		}
		else
		{
			if (w == 2)
			{
				if (rank[f1] < rank[f2])
				{
					id[f1] = f2;
				}
				else
				{
					if (rank[f2] < rank[f1])
						id[f2] = f1;
					else
					{
						id[f1] = f2;
						rank[f2]++;
					}
				}
				unionPointers++;
			}
			else
				id[f1] = f2;
		}
			
		
		unionPointers++;
	
		if (worstCaseUnion < unionPointers - oldPointers)
		{
			worstCaseUnion = unionPointers - oldPointers;
		}
	}
	
	public boolean find(int x, int y)
	{
		finds++;
		int oldPointers = findPointers;
		
		int f1 = find(x);
		int f2 = find(y);
		
		if (worstCaseFind < findPointers - oldPointers)
		{
			worstCaseFind = findPointers - oldPointers;
		}
		
		return f1 == f2;
	}

	public int find(int x)
	{
		
		
		if (pc == 1) { while (x != id[x]){ id[x] = id[id[x]]; findPointers++; x = id[x]; findPointers++;} }
		if (pc == 2) { if (x != id[x]) {id[x] = find(id[x]); findPointers++; } return id[x];}
		if (pc == 0) { while (x != id[x]) {x = id[x]; findPointers++; } }
		
		return x;
	}
	
	public int findU(int x)
	{
		if (pc == 1) { while (x != id[x]){ id[x] = id[id[x]]; unionPointers++; x = id[x]; unionPointers++;} }
		if (pc == 2) { if (x != id[x]) {id[x] = findU(id[x]); unionPointers++; } return id[x];}
		if (pc == 0) { while (x != id[x]) {x = id[x]; unionPointers++; } }
		
		return x;
	}
	
	public int getFindPointers()
	{
		return findPointers;
	}
	
	public int getUnionPointers()
	{
		return unionPointers;
	}
	
	public int getWorstCaseFind()
	{
		return worstCaseFind;
	}
	
	public int getWorstCaseUnion()
	{
		return worstCaseUnion;
	}
	
	public int getUnions()
	{
		return unions;
	}
	
	public int getFinds()
	{
		return finds;
	}
	
	public static void main(String[] args)
	{
		SuperUnionFind suf = new SuperUnionFind(10, 2, 2);
		System.out.println("1: " + suf.getUnionPointers());
		suf.union(3, 4);
		System.out.println("2: " + suf.getUnionPointers());
		suf.union(4, 9);
		System.out.println("3: " + suf.getUnionPointers());
		suf.union(8, 0);
		System.out.println("4: " + suf.getUnionPointers());
		suf.union(2, 3);
		System.out.println("5: " + suf.getUnionPointers());
		suf.union(5, 6);
		System.out.println("6: " + suf.getUnionPointers());
		suf.union(5, 9);
		System.out.println("7: " + suf.getUnionPointers());
		suf.union(7, 3);
		System.out.println("8: " + suf.getUnionPointers());
		suf.union(4,8);
		System.out.println("9: " + suf.getUnionPointers());
		System.out.println(suf.find(0, 1));
		System.out.println(suf.find(6, 8));
		
	}
}
