#include <config.h>
#include <graph/MixtureNode.h>

#include <utility>
#include <vector>
#include <cfloat>
#include <stdexcept>

using std::pair;
using std::vector;
using std::map;
using std::invalid_argument;
using std::logic_error;

MixtureNode::MixtureNode(vector<Node *> const &index,
			 vector<pair<Index, Node *> > const &parameters)
  : DeterministicNode(parameters[0].second->data.dim(true)), _index(index)
{
  for (vector<Node *>::const_iterator i = index.begin(); i != index.end(); ++i)
    {
      if ((*i)->data.length() != 1 || !(*i)->data.isDiscreteValued()) {
	throw invalid_argument("Invalid index parameter for mixture node");
      }
      this->addParent(*i);	
    }

  unsigned int ndim = parameters.size();
  Index const &default_dim = data.dim(false);
  bool isdiscrete = true;
  for (unsigned int i = 0; i < ndim; ++i) {
    Node *node = parameters[i].second;
    if (!node) {
      throw invalid_argument("Null parameter in MixtureNode");
    }
    if (node->data.dim(true) != default_dim) {
      throw invalid_argument("Dimension mismatch for MixtureNode parameters");
    }
    if (!node->data.isDiscreteValued()) {
      isdiscrete = false;
    }
    this->addParent(node);
    _map[parameters[i].first] = node;
  }
  data.setDiscreteValued(isdiscrete);
}

MixtureNode::~MixtureNode()
{
}

void MixtureNode::forwardSample()
{
  Index i(_index.size());
  for (unsigned int j = 0; j < _index.size(); ++j) {
    i[j] = static_cast<long>(_index[j]->data.value()[0]);
  }
  map<Index,Node*>::iterator p = _map.find(i);
  if (p != _map.end()) {
    data.setValue(p->second->data.value(), 
		  p->second->data.length());
  }
  else {
    throw logic_error("Invalid index in MixtureNode");
  }
}

vector<Node*> const &MixtureNode::index() const
{
  return _index;
}

MixtureNode const *asMixture(Node const *node)
{
  return dynamic_cast<MixtureNode const*>(node);
}

bool isMixture(Node const *node)
{
  return dynamic_cast<MixtureNode const*>(node);
}


syntax highlighted by Code2HTML, v. 0.9.1