package inference.axioms;

import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;

import exceptions.SubstituteFailedException;
import inference.EquationAxiom;
import models.Position;
import models.algebra.Variable;
import models.formulas.Formula;
import models.formulas.meta.MetaDependencyFormula;
import models.formulas.meta.MetaEquationFormula;
import models.terms.EvaluatableTerm;
import models.terms.meta.MatchConstraint;
import models.terms.meta.MetaDynamicDependency;
import models.terms.meta.MetaDynamicDependencyTerm;
import models.terms.meta.MetaEvaluatableTermVariable;
import models.terms.meta.MetaRDLTerm;
import models.terms.meta.MetaTermGenerator;

public class MapComposition extends EquationAxiom {
	
	public MapComposition() {
		super("Map Composition");
		assumptions.add(new MetaDependencyFormula(
				new MetaDynamicDependency(
						new MetaTermGenerator() {
							@Override
							public MetaRDLTerm generate(int curIndex, int curDepth, int maxIndex, int maxDepth, Map<Object, Object> context) {
								curIndex -= 2;
								context.put("tIndex", curIndex + 1);
								return new MetaEvaluatableTermVariable(new Variable("ve" + curIndex));
							}
						},
						new MetaEvaluatableTermVariable(new Variable("se")),
						new MetaEvaluatableTermVariable(new Variable("te"))
				)
		));
		
		assumptions.add(new MetaDependencyFormula(
				new MetaDynamicDependency(
						new MetaTermGenerator() {
							@Override
							public MetaRDLTerm generate(int curIndex, int curDepth, int maxIndex, int maxDepth, Map<Object, Object> context) {
								curIndex -= 1;
								context.put("uIndex", curIndex + 1);
								return new MetaEvaluatableTermVariable(new Variable("ue" + curIndex));
							}
						},
						new MetaEvaluatableTermVariable(new Variable("te"))
				)
		));
		
		conclusion = new MetaEquationFormula(
				new MetaDynamicDependencyTerm(
						new MetaTermGenerator() {
							@Override
							public MetaRDLTerm generate(int curIndex, int curDepth, int maxIndex, int maxDepth, Map<Object, Object> context) {
								curIndex -= 3;
								int tIndex = (Integer) context.get("tIndex");
								if (curIndex >= tIndex * 2) {
									return null;
								}
								if (curIndex % 2 == 0) {
									return new MetaEvaluatableTermVariable(new Variable("ve" + curIndex / 2));
								}
								return new MetaEvaluatableTermVariable(new Variable("v" + curIndex / 2));
							}
						},
						new MetaEvaluatableTermVariable(new Variable("se")),
						new MetaEvaluatableTermVariable(new Variable("te")),
						new MetaDynamicDependencyTerm(
								new MetaTermGenerator() {
									@Override
									public MetaRDLTerm generate(int curIndex, int curDepth, int maxIndex, int maxDepth, Map<Object, Object> context) {
										curIndex -= 1;
										int uIndex = (Integer) context.get("uIndex");
										if (curIndex >= uIndex * 2) {
											return null;
										}
										if (curIndex % 2 == 0) {
											return new MetaEvaluatableTermVariable(new Variable("ue" + curIndex / 2));
										}
										return new MetaEvaluatableTermVariable(new Variable("u" + curIndex / 2));
									}
								},
								new MetaEvaluatableTermVariable(new Variable("te"))
						)
				),
				new MetaDynamicDependencyTerm(
						new MetaTermGenerator() {
							@Override
							public MetaRDLTerm generate(int curIndex, int curDepth, int maxIndex, int maxDepth, Map<Object, Object> context) {
								curIndex -= 1;
								int uIndex = (Integer) context.get("uIndex");
								int tIndex = (Integer) context.get("tIndex");
								if (curIndex >= uIndex * 2 + tIndex * 2) {
									return null;
								}
								if (curIndex % 2 == 0 && curIndex < uIndex * 2) {
									return new MetaEvaluatableTermVariable(new Variable("ue" + curIndex / 2));
								} else if (curIndex < uIndex * 2) {
									return new MetaEvaluatableTermVariable(new Variable("u" + curIndex / 2));
								}
								curIndex -= uIndex * 2;
								if (curIndex % 2 == 0) {
									return new MetaEvaluatableTermVariable(new Variable("ve" + curIndex / 2));
								}
								return new MetaEvaluatableTermVariable(new Variable("v" + curIndex / 2));
							}
						},
						new MetaEvaluatableTermVariable(new Variable("se"))
				)
		);
	}

	@Override
	public Set<EvaluatableTerm> apply(List<Formula>assumptions,  EvaluatableTerm term, MatchConstraint constraint) {
		Set<EvaluatableTerm> result = new HashSet<>();
		if (assumptions.size() < this.assumptions.size()) {
			return new HashSet<>();
		}
		Set<MatchConstraint> matchResult = assumptionMatch(assumptions, constraint);
		
		
		Set<MatchConstraint> conclusionLeftMatchResult = conclusionLeftSideHandMatch(term, matchResult);
		for (MatchConstraint leftConst: conclusionLeftMatchResult) {
			leftConst.getContext().put("isLeft", true);
		}
		Set<MatchConstraint> conclusionRightMatchResult = conclusionRightSideHandMatch(term, matchResult);
		for (MatchConstraint rightConst: conclusionRightMatchResult) {
			rightConst.getContext().put("isLeft", false);
		}
		Set<MatchConstraint> conclusionMatchResult = new HashSet<>();
		conclusionMatchResult.addAll(conclusionLeftMatchResult);
		conclusionMatchResult.addAll(conclusionRightMatchResult);
		
		MetaEquationFormula metaConclusion = (MetaEquationFormula) this.conclusion;
		
		for (MatchConstraint matchRes: conclusionMatchResult) {
			try {
				boolean isLeft = (Boolean) matchRes.getContext().get("isLeft");
				if (isLeft) {
					int n = (Integer) matchRes.getContext().get(new Position());
					int k = (Integer) matchRes.getContext().get(new Position(0, 2));
					matchRes.getContext().put(new Position(), n + k - 3);
					EvaluatableTerm rightSide = (EvaluatableTerm) ((MetaRDLTerm) metaConclusion.getRightSideHand()).substitute(matchRes.getBinding(), matchRes.getContext());
					result.add(rightSide);
				} else {
					int n = (Integer) matchRes.getContext().get(new Position());
					int k = (Integer) matchRes.getContext().get("uIndex") * 2;
					matchRes.getContext().put(new Position(), n - k + 2);
					matchRes.getContext().put(new Position(0, 2), k + 1);
					EvaluatableTerm leftSide = (EvaluatableTerm) ((MetaRDLTerm) metaConclusion.getLeftSideHand()).substitute(matchRes.getBinding(), matchRes.getContext());
					result.add(leftSide);
				}
			} catch (SubstituteFailedException e) {
				continue;
			}
		}
		return result;
	}
	
}
