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.EvaluatableTerm;
import models.terms.PrimedTerm;
import models.terms.Resource;

public class RewriteInferenceSystem {

	List<EquationFormula> constraintFormulas = new ArrayList<>();
	List<EquationFormula> invariantFormulas = new ArrayList<>();
	EquationFormula inputFormula;
	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)) {
				inputFormula = (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: " + inputFormula.toString());
		System.out.println("conditionalFormulas: " + conditionalFormulas.toString());
		System.out.println("conclusion: " + conclusion.toString());
	}
	
	public boolean inference() {
		ResourceTree baseTree = expandTree();
		System.out.println("baseTree: " + baseTree);
		Set<ResourceTree> result = rewriteTree(baseTree);
		for (ResourceTree tree: result) {
			System.out.println("=====================================");
			tree.debugAllPath();
		}
		return false;
	}
	
	
	private ResourceTree expandTree() {
		ResourceTree inputResourceTree = new ResourceTree(inputFormula.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, Resource> 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 Set<ResourceTree> rewriteTree(ResourceTree expandedInputResourceTree) {
		Set<ResourceTree> result = new HashSet<>();
		result.add(expandedInputResourceTree);
		Map<ResourceTree, ResourceTree> rewritable = new HashMap<>();
		for (EquationFormula formula : constraintFormulas) {
			rewritable.put(new ResourceTree(formula.getRightSideHand()), new ResourceTree(formula.getLeftSideHand()));
		}
		for (EquationFormula formula : conditionalFormulas) {
			rewritable.put(new ResourceTree(formula.getRightSideHand()), new ResourceTree(formula.getLeftSideHand()));
		}
		rewritable.put(new ResourceTree(inputFormula.getLeftSideHand()), new ResourceTree(inputFormula.getRightSideHand()));
		
		Deque<ResourceTree> treeQueue = new ArrayDeque<>();
		treeQueue.add(expandedInputResourceTree);
		while (! treeQueue.isEmpty()) {
			ResourceTree curBaseTree = treeQueue.pollFirst();
			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);
				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, Resource> 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 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;
	}
	
	
}
