package inference.rewrite;

import java.util.ArrayDeque;
import java.util.ArrayList;
import java.util.Deque;
import java.util.HashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;

import models.formulas.EquationFormula;
import models.formulas.Formula;
import models.formulas.Then;
import models.terms.DependencyTerm;
import models.terms.EvaluatableTerm;
import models.terms.PrimedTerm;
import models.terms.Resource;

public class RewriteInferenceSystem {

	List<EquationFormula> constraintFormulas = new ArrayList<>();
	List<EquationFormula> invariantFormulas = new ArrayList<>();
	List<EquationFormula> inputFormulas = new ArrayList<>();
	List<EquationFormula> conditionalFormulas = new ArrayList<>();
	List<Formula> otherFormulas = new ArrayList<>();
	EquationFormula conclusion;
	
	
	public RewriteInferenceSystem(List<Formula> assumptions, Formula conclusion) {
		for (Formula assumption : assumptions) {
			if (inputFormulaCheck(assumption)) {
				inputFormulas.add((EquationFormula) assumption);
			} else if(invariantFormulaCheck(assumption)) {
				invariantFormulas.add((EquationFormula) assumption);
			} else if (assumption instanceof EquationFormula) {
				constraintFormulas.add((EquationFormula) assumption);
			} else {
				otherFormulas.add(assumption);
			}
		}
		
		if (conclusion instanceof Then then) {
			conditionalFormulas.add((EquationFormula) then.getCondition());
			this.conclusion = (EquationFormula) then.getResult();
		} else {
			this.conclusion = (EquationFormula) conclusion;
		}
		
	}
	
	public void debug() {
		System.out.println("constraintFormulas: " + constraintFormulas.toString());
		System.out.println("invariantFormulas: " + invariantFormulas.toString());
		System.out.println("inputFormula: " + inputFormulas.toString());
		System.out.println("conditionalFormulas: " + conditionalFormulas.toString());
		System.out.println("conclusion: " + conclusion.toString());
	}
	
	public boolean inference() {
		EquationFormula defCheck = definitonFormulaCheck();
		if (defCheck != null) {
			System.out.println("non definition formula is detected. " + defCheck.toString() );
			return false;
		}
		if (!loopCheck()) {
			System.out.println("Loop is detected");
			return false;
		}
		for (EquationFormula inputFormula : inputFormulas) {
			System.out.println("---------------------------" + inputFormula + "-------------------------------------------------");
//			ResourceTree baseTree = expandTree(i);
//			System.out.println("baseTree: " + baseTree);
//			Set<ResourceTree> result = rewriteTree(baseTree, i);
//			System.out.println("=================result====================");
//			for (ResourceTree tree: result) {
//				tree.debugAllPath();
//			}
//			if (result.contains(new ResourceTree(conclusion.getLeftSideHand())) && result.contains(new ResourceTree(conclusion.getRightSideHand()))) {
//				System.out.println("yes");
//			} else {
//				System.out.println("no");
//			}
//			System.out.println("=====================================");
			Map<ResourceTree, Set<RewriteGraphNode>> leftRewriteGraph = new HashMap<>();
			Map<ResourceTree, Set<RewriteGraphNode>> rightRewriteGraph = new HashMap<>();
			Set<ResourceTree> conclusionLeftResult = rewriteTree(new ResourceTree(conclusion.getLeftSideHand()), inputFormula, leftRewriteGraph);
			Set<ResourceTree> conclusionRightResult = rewriteTree(new ResourceTree(conclusion.getRightSideHand()), inputFormula, rightRewriteGraph);
			conclusionLeftResult.retainAll(conclusionRightResult);
			ResourceTree resultRoot = conclusionLeftResult.iterator().next();
			showRewriteGraph(resultRoot, leftRewriteGraph);
			System.out.println("=============================================================----");
			showRewriteGraph(resultRoot, rightRewriteGraph);
//			System.out.println(conclusionLeftResult);
			System.out.println(conclusionLeftResult.size() != 0);
			
		}
		
		return false;
	}
	
	
	private ResourceTree expandTree(int inputFormulaIndex) {
		ResourceTree inputResourceTree = new ResourceTree(inputFormulas.get(inputFormulaIndex).getLeftSideHand());
		List<ResourceTree> constraintResourceTree =	constraintFormulas.stream().map(v -> new ResourceTree(v.getRightSideHand())).toList();
		List<ResourceTree> conditionalResourceTree = conditionalFormulas.stream().map(v -> new ResourceTree(v.getLeftSideHand())).toList();
		boolean rewrited = true;
		while (rewrited) {
			rewrited = false;
			inputResourceTree.debugAllPath();
			for (ResourceTree tree: constraintResourceTree) {
				Position matchPos = treeJoinCheck(inputResourceTree, tree);
				if (matchPos != null) {
					inputResourceTree = joinTree(inputResourceTree, matchPos, tree);
					rewrited = true;
				} else {
					matchPos = treeJoinCheck(tree, inputResourceTree);
					if (matchPos != null) {
						inputResourceTree = joinTree(tree, matchPos, inputResourceTree);
						rewrited = true;
					}
				}
			}
			for (ResourceTree tree: conditionalResourceTree) {
				Position matchPos = treeJoinCheck(inputResourceTree, tree);
				if (matchPos != null) {
					inputResourceTree = joinTree(inputResourceTree, matchPos, tree);
					rewrited = true;
				} else {
					matchPos = treeJoinCheck(tree, inputResourceTree);
					if (matchPos != null) {
						inputResourceTree = joinTree(tree, matchPos, inputResourceTree);
						rewrited = true;
					}
				}
			}
		}
		return inputResourceTree;
	}
	
	private Position treeJoinCheck(ResourceTree leftTree, ResourceTree rightTree) {
		Deque<Position> positionQueue = new ArrayDeque<>();
		for (Position nextPos: leftTree.getChildren(new Position())) {
			positionQueue.add(nextPos);
		}
		while(! positionQueue.isEmpty()) {
			Position curPos = positionQueue.pollFirst();
			if (treeJoinCheck(leftTree, curPos, rightTree, new Position())) {
				return curPos;
			}
			for (Position nextPos: leftTree.getChildren(curPos)) {
				positionQueue.add(nextPos);
			}
		}
		return null;
	}
	
	private boolean treeJoinCheck(ResourceTree leftTree, Position leftTreePos, ResourceTree rightTree, Position rightTreePos) {
		if (leftTree.getResource(leftTreePos).equals(rightTree.getResource(rightTreePos))) {
			boolean result = true;
			Set<Position> used = new HashSet<>();
			for (Position nextLeftPos: leftTree.getChildren(leftTreePos)) {
				boolean someTreeMatched = false;
				for (Position nextRightPos: rightTree.getChildren(rightTreePos)) {
					if (treeJoinCheck(leftTree, nextLeftPos, rightTree, nextRightPos)) {
						used.add(nextRightPos);
						someTreeMatched = true;
						break;
					}
				}
				result &= someTreeMatched;
			}
			return result;
		}
		return false;
	}
	
	private record PositionPair(Position resultPos, Position rightPos) {};
	
	private ResourceTree joinTree(ResourceTree leftTree, Position joinPos, ResourceTree rightTree) {
		Map<Position, List<Position>> tree = new HashMap<>();
		Map<Position, EvaluatableTerm> resourceMap = new HashMap<>();
		
		Deque<PositionPair> positionQue = new ArrayDeque<>();
		Position rootPos = new Position();
		positionQue.add(new PositionPair(rootPos, null));
		tree.put(rootPos, new ArrayList<>());
		resourceMap.put(rootPos, leftTree.getResource(rootPos));
		while (! positionQue.isEmpty()) {
			PositionPair curPosPair = positionQue.pollFirst();
			Position curPos = curPosPair.resultPos();
			Position curRightPos = curPosPair.rightPos();
			if (curRightPos == null) {
				for (Position nextPos: leftTree.getChildren(curPos)) {
					tree.get(curPos).add(nextPos);
					tree.put(nextPos, new ArrayList<>());
					resourceMap.put(nextPos, leftTree.getResource(nextPos));
					if (nextPos.startWith(joinPos)) {
						positionQue.add(new PositionPair(nextPos, new Position()));
					} else {
						positionQue.add(new PositionPair(nextPos, null));
					}
				}
			} else {
				for (int i = 0; i < rightTree.getChildren(curRightPos).size(); i++) {
					Position nextRightPos = rightTree.getChildren(curRightPos).get(i);
					Position nextResultPos = curPos.addPath(i);
					tree.get(curPos).add(nextResultPos);
					tree.put(nextResultPos, new ArrayList<>());
					resourceMap.put(nextResultPos, rightTree.getResource(nextRightPos));
					positionQue.add(new PositionPair(nextResultPos, nextRightPos));
				}
			}
		}
		
		return new ResourceTree(tree, resourceMap);
		
	}
	
	private void showRewriteGraph(ResourceTree root, Map<ResourceTree, Set<RewriteGraphNode>> rewriteGraph) {
		List<Set<ResourceTree>> result = new ArrayList<>();
		result.add(new HashSet<>());
		result.get(0).add(root);
		boolean isChanged = true;
		while (isChanged) {
			Set<ResourceTree> nextTrees = new HashSet<>();
			isChanged = false;
			for (ResourceTree curTree : result.get(result.size() - 1)) {
				if (rewriteGraph.containsKey(curTree)) {
					isChanged |= !rewriteGraph.get(curTree).isEmpty();
					for (RewriteGraphNode node: rewriteGraph.get(curTree)) {
						nextTrees.add(node.baseTree());
					}
				}
			}
			result.add(nextTrees);
		}
		result.remove(result.size() - 1);
		for (int i = result.size() - 1; i >= 0; i--) {
			Set<ResourceTree> curTrees = result.get(i);
			StringBuilder sb = new StringBuilder();
			for (ResourceTree rt : curTrees) {
				sb.append(rewriteGraph.get(rt));
				sb.append(", \n");
			}
			sb.delete(sb.length() - 1, sb.length());
			System.out.println(sb.toString());
			System.out.println("--------------------------");
		}
		System.out.println(root);
	}
	
	private record RewriteGraphNode(ResourceTree baseTree, ResourceTree leftSideHand, ResourceTree rightSideHand, ResourceTree result) {
		@Override
		public String toString() {
			return baseTree.toString() + " --( " + leftSideHand.toString() + " = " + rightSideHand.toString() + " )--> " + result ;
		}
		@Override
		public int hashCode() {
			return toString().hashCode();
		}
		@Override
		public boolean equals(Object another) {
			if ( ! (another instanceof RewriteGraphNode)) {
				return false;
			}
			RewriteGraphNode node = (RewriteGraphNode) another;
			return this.baseTree.equals(node.baseTree()) && this.leftSideHand.equals(node.leftSideHand()) && this.rightSideHand.equals(node.rightSideHand()) && this.result.equals(node.result());
		}
	};
	
	private Set<ResourceTree> rewriteTree(ResourceTree expandedInputResourceTree, EquationFormula inputFormula, Map<ResourceTree, Set<RewriteGraphNode>> rewriteGraph) {
		Set<ResourceTree> result = new HashSet<>();
		Map<ResourceTree, ResourceTree> rewritable = new HashMap<>();
		for (EquationFormula formula : constraintFormulas) {
			rewritable.put(new ResourceTree(formula.getLeftSideHand()), new ResourceTree(formula.getRightSideHand()));
		}
		for (EquationFormula formula : conditionalFormulas) {
			rewritable.put(new ResourceTree(formula.getLeftSideHand()), new ResourceTree(formula.getRightSideHand()));
		}
		for (EquationFormula formula : invariantFormulas) {
			rewritable.put(new ResourceTree(formula.getLeftSideHand()), new ResourceTree(formula.getRightSideHand()));
		}
		rewritable.put(new ResourceTree(inputFormula.getLeftSideHand()), new ResourceTree(inputFormula.getRightSideHand()));
		
		Deque<ResourceTree> treeQueue = new ArrayDeque<>();
		treeQueue.add(expandedInputResourceTree);
		Set<RewriteGraphNode> used = new HashSet<>();
		while (! treeQueue.isEmpty()) {
			ResourceTree curBaseTree = treeQueue.pollFirst();
			if (result.contains(curBaseTree)) {
				continue;
			}
			result.add(curBaseTree);
			for (ResourceTree from: rewritable.keySet()) {
				ResourceTree to = rewritable.get(from);
				Set<Position> matchPositions = treeMatchCheck(curBaseTree, from);
				if (matchPositions == null) {
					continue;
				}
				ResourceTree res = rewrite(curBaseTree, matchPositions, to);
				if (! rewriteGraph.containsKey(res)) {
					rewriteGraph.put(res, new HashSet<>());
				}
				RewriteGraphNode nextNode = new RewriteGraphNode(curBaseTree, from, to, res);
				if (!used.contains(nextNode)) {
					rewriteGraph.get(res).add(nextNode);
					used.add(nextNode);
				}
				treeQueue.add(res);
			}
		}
		
		return result;
	}
	
	private Set<Position> treeMatchCheck(ResourceTree baseTree, ResourceTree matchTree) {
		Deque<Position> positionQueue = new ArrayDeque<>();
		positionQueue.add(new Position());
		while (! positionQueue.isEmpty()) {
			Position curPos = positionQueue.pollFirst();
			Set<Position> result = treeMatchCheck(baseTree, curPos, matchTree);
			if (result != null) {
				return result;
			}
			for (Position child: baseTree.getChildren(curPos)) {
				positionQueue.add(child);
			}
		}
		return null;
	}
	
	private Set<Position> treeMatchCheck(ResourceTree baseTree, Position startPosition, ResourceTree matchTree) {
		Set<Position> result = new HashSet<>();
		Deque<Position> basePosQueue = new ArrayDeque<>();
		Deque<Position> matchPosQueue = new ArrayDeque<>();
		basePosQueue.add(startPosition);
		matchPosQueue.add(new Position());
		if (! baseTree.getResource(startPosition).equals(matchTree.getResource(new Position()))) {
			return null;
		}
		while (! basePosQueue.isEmpty()) {
			Position curBasePos = basePosQueue.pollFirst();
			Position curMatchPos = matchPosQueue.pollFirst();
			result.add(curBasePos);
			if (matchTree.getChildren(curMatchPos).size() == 0) {
				continue;
			}
			Set<Position> used = new HashSet<>(); 
			for (Position nextBasePos: baseTree.getChildren(curBasePos)) {
				for (Position nextMatchPos: matchTree.getChildren(curMatchPos)) {
					if (used.contains(nextMatchPos)) continue;
					if (baseTree.getResource(nextBasePos).equals(matchTree.getResource(nextMatchPos))) {
						used.add(nextMatchPos);
						basePosQueue.add(nextBasePos);
						matchPosQueue.add(nextMatchPos);
						break;
					}
					return null;
				}
			}
		}
		return result;
	}
	
	private ResourceTree rewrite(ResourceTree baseTree, Set<Position> matchPositions, ResourceTree toTree) {
		Map<Position, List<Position>> resultTree = new HashMap<>();
		Map<Position, EvaluatableTerm> resultResourceMap = new HashMap<>();
		Position resultCurPos = new Position();
		Position baseCurPos = new Position();
		Position toCurPos = new Position();
		Deque<Position> resultPosQueue = new ArrayDeque<>();
		Deque<Position> basePosQueue = new ArrayDeque<>();
		Deque<Position> toPosQueue = new ArrayDeque<>();
		
		resultPosQueue.add(resultCurPos);
		basePosQueue.add(baseCurPos);
		toPosQueue.add(toCurPos);
		
		
		while (! basePosQueue.isEmpty()) {
			baseCurPos = basePosQueue.pollFirst();
			if (matchPositions.contains(baseCurPos)) {
				List<Position> lastPositions = new ArrayList<>();
				Deque<Position> matchPosQueue = new ArrayDeque<>();
				matchPosQueue.add(baseCurPos);
				while (! matchPosQueue.isEmpty()) {
					Position matchPos = matchPosQueue.pollFirst();
					for (Position nextMatchPos: baseTree.getChildren(matchPos)) {
						if (matchPositions.contains(nextMatchPos)) {
							matchPosQueue.add(nextMatchPos);
						} else {
							lastPositions.add(nextMatchPos);
						}
					}
				}
				List<Position> resultLastPositions = new ArrayList<>();
				while (! toPosQueue.isEmpty()) {
					toCurPos = toPosQueue.pollFirst();
					resultCurPos = resultPosQueue.pollFirst();
					resultTree.put(resultCurPos, new ArrayList<>());
					resultResourceMap.put(resultCurPos, toTree.getResource(toCurPos));
					if (toTree.getChildren(toCurPos).size() == 0) {
						resultLastPositions.add(resultCurPos);
						continue;
					}
					for (int i = 0; i < toTree.getChildren(toCurPos).size(); i++) {
						Position resultNextPos = resultCurPos.addPath(i);
						resultTree.get(resultCurPos).add(resultNextPos);
						resultPosQueue.add(resultNextPos);
						toPosQueue.add(toCurPos.addPath(i));
					}
				}
				
				for (Position resultLastPos: resultLastPositions) {
					Deque<Position> resultLastPosQueue = new ArrayDeque<>();
					Deque<Position> lastPosQueue = new ArrayDeque<>();
					for (int i = 0; i < lastPositions.size(); i++) {
						Position nextResultLastPos = resultLastPos.addPath(i);
						resultTree.get(resultLastPos).add(nextResultLastPos);
						resultLastPosQueue.add(nextResultLastPos);
						lastPosQueue.add(lastPositions.get(i));
					}
					while (! resultLastPosQueue.isEmpty()) {
						Position curResultLastPos = resultLastPosQueue.pollFirst();
						Position curLastPos = lastPosQueue.pollFirst();
						resultTree.put(curResultLastPos, new ArrayList<>());
						resultResourceMap.put(curResultLastPos, baseTree.getResource(curLastPos));
						for (int i = 0; i < baseTree.getChildren(curLastPos).size(); i++) {
							Position nextResultLastPos = curResultLastPos.addPath(i);
							resultTree.get(curResultLastPos).add(nextResultLastPos);
							resultLastPosQueue.add(nextResultLastPos);
							lastPosQueue.add(curLastPos.addPath(i));
						}
					}
				}
			} else {
				resultCurPos = resultPosQueue.pollFirst();
				resultTree.put(resultCurPos, new ArrayList<>());
				resultResourceMap.put(resultCurPos, baseTree.getResource(baseCurPos));
				for (int i = 0; i < baseTree.getChildren(baseCurPos).size(); i++) {
					Position resultNextPos = resultCurPos.addPath(i);
					resultTree.get(resultCurPos).add(resultNextPos);
					resultPosQueue.add(resultNextPos);
					basePosQueue.add(baseCurPos.addPath(i));
				}
			}
		}
		
		return new ResourceTree(resultTree, resultResourceMap);
	} 
	
	private boolean loopCheck() {
		Map<EvaluatableTerm, EvaluatableTerm> equationGraph = new HashMap<>();
		for (EquationFormula formula : constraintFormulas) {
			equationGraph.put(formula.getLeftSideHand(), formula.getRightSideHand());
		}
		for (EquationFormula formula : invariantFormulas) {
			equationGraph.put(formula.getLeftSideHand(), formula.getRightSideHand());
		}
		for (EquationFormula inputFormula : inputFormulas) {
			equationGraph.put(inputFormula.getLeftSideHand(), inputFormula.getRightSideHand());
		}
		for (EquationFormula formula : conditionalFormulas) {
			equationGraph.put(formula.getLeftSideHand(), formula.getRightSideHand());
		}
		for (EvaluatableTerm leftTerm : equationGraph.keySet()) {
			Set<EvaluatableTerm> used = new HashSet<>();
			used.add(leftTerm);
			EvaluatableTerm rightTerm = equationGraph.get(leftTerm);
			if (checkUsedTerms(used, rightTerm)) {
				return false;
			}
			while (equationGraph.containsKey(rightTerm)) {
				leftTerm = rightTerm;
				used.add(leftTerm);
				rightTerm = equationGraph.get(leftTerm);
				if (checkUsedTerms(used, rightTerm)) {
					return false;
				}
			}
		}
		return true;
	}
	
	private boolean checkUsedTerms(Set<EvaluatableTerm> used, EvaluatableTerm term) {
		if (used.contains(term)) {
			return true;
		}
		if (term instanceof DependencyTerm depTerm) {
			if (checkUsedTerms(used, depTerm.getDependingTerm())) {
				return true;
			}
			for (EvaluatableTerm te : depTerm.getDependedTerms()) {
				if (checkUsedTerms(used, te)) {
					return true;
				}
			}
			for (EvaluatableTerm te : depTerm.getArgumentTerms()) {
				if (checkUsedTerms(used, te)) {
					return true;
				}
			}
		}
		return false;
	}
	
	private EquationFormula definitonFormulaCheck() {
		for (EquationFormula eq : constraintFormulas) {
			if (!definitionFormulaCheck(eq)) {
				return eq;
			}
		}
		for (EquationFormula inputFormula : inputFormulas) {
			if (!definitionFormulaCheck(inputFormula)) {
				return inputFormula;
			}
		}
		for (EquationFormula eq : invariantFormulas) {
			if (!definitionFormulaCheck(eq)) {
				return eq;
			}
		}
		for (EquationFormula eq : conditionalFormulas) {
			if (!definitionFormulaCheck(eq)) {
				return eq;
			}
		}
		return null;
	}
	
	private boolean definitionFormulaCheck(Formula formula) {
		if (formula instanceof EquationFormula equation) {
			return singleTermCheck(equation.getLeftSideHand());
		} else if (formula instanceof Then then) {
			Formula f1 = then.getCondition();
			Formula f2 = then.getResult();
			if (f1 instanceof EquationFormula eq1 && f2 instanceof EquationFormula eq2) {
				return singleTermCheck(eq1.getLeftSideHand()) && singleTermCheck(eq2.getLeftSideHand());
			}
		}
		return false;
	}
	
	private boolean singleTermCheck(EvaluatableTerm term) {
		if (term instanceof PrimedTerm primedTerm) {
			return singleTermCheck(primedTerm.getPrimedTerm());
		}
		if (term instanceof Resource) {
			return true;
		}
		DependencyTerm depTerm = (DependencyTerm) term;
		if (!(depTerm.getDependingTerm() instanceof Resource)) {
			if (depTerm.getDependingTerm() instanceof PrimedTerm primedTerm) {
				if (! (primedTerm.getPrimedTerm() instanceof Resource)) {
					return false;
				}
			} else {
				return false;
			}
		}
		for (EvaluatableTerm te : depTerm.getDependedTerms()) {
			if (!(te instanceof Resource)) {
				if (te instanceof PrimedTerm primedTerm) {
					if (! (primedTerm.getPrimedTerm() instanceof Resource)) {
						return false;
					}
				} else {
					return false;
				}
			}
			
		}
		for (EvaluatableTerm te : depTerm.getArgumentTerms()) {
			if (!(te instanceof Resource)) {
				if (te instanceof PrimedTerm primedTerm) {
					if (! (primedTerm.getPrimedTerm() instanceof Resource)) {
						return false;
					}
				} else {
					return false;
				}
			}
		}
		return true;
	}
	
	private boolean inputFormulaCheck(Formula formula) {
		if (! (formula instanceof EquationFormula equation)) {
			return false;
		}
		List<Resource> resources  = new ArrayList<>(equation.getLeftSideHand().getSubTerms(Resource.class).values());
		resources.addAll(equation.getRightSideHand().getSubTerms(Resource.class).values());
		for (Resource resource : resources) {
			if (resource.getOrder() == 0) {
				return true;
			}
		}
		return false;
	}
	
	private boolean invariantFormulaCheck(Formula formula) {
		if (! (formula instanceof EquationFormula equation)) {
			return false;
		}
		EvaluatableTerm leftSideHand = equation.getLeftSideHand();
		EvaluatableTerm rightSideHand = equation.getRightSideHand();
		
		if (leftSideHand instanceof PrimedTerm) {
			return ((PrimedTerm) leftSideHand).getPrimedTerm().equals(rightSideHand);
		} else if (rightSideHand instanceof PrimedTerm) {
			return ((PrimedTerm) rightSideHand).getPrimedTerm().equals(leftSideHand);
		}
		
		return false;
	}
	
	
	
}
