require 'minitest/autorun'
require_relative 'example'
class StudentTests < Minitest::Test
  def setup
    @rows=ExamExperiment.load_rows(File.join(__dir__,'students.csv'))
    @training,@held=@rows.partition { |r| r['cohort']=='train' }
  end
  def test_fit_only_training_and_do_not_refit_on_prediction
    fitted=ExamExperiment.fit(@training);before=fitted[0].mean_vec.to_a
    expected=ExamExperiment::FEATURES.map { |key| @training.sum { |r| Float(r[key]) }/@training.length }
    assert_equal expected,before
    ExamExperiment.probabilities(fitted,@held)
    assert_equal before,fitted[0].mean_vec.to_a
    assert_raises(ArgumentError) { ExamExperiment.fit(@rows) }
  end
  def test_probability_class_and_threshold
    fit=ExamExperiment.fit(@training);p=ExamExperiment.probabilities(fit,@held)
    assert_equal [0,1],fit[1].classes.to_a
    assert p.all? { |v| v.finite? && v.between?(0,1) }
    low=ExamExperiment.confusion([0,1],[0.6,0.9],0.5)
    high=ExamExperiment.confusion([0,1],[0.6,0.9],0.8)
    assert_equal({tp:1,fp:1,tn:0,fn:0},low)
    assert_equal({tp:1,fp:0,tn:1,fn:0},high)
  end
  def test_reject_bad_features
    ['',nil,'unknown','NaN','Infinity','-1','101'].each do |value|
      assert_raises(ArgumentError) { ExamExperiment.features({'attendance_pct'=>value,'practice_hours'=>'4'}) }
    end
    assert_raises(ArgumentError) { ExamExperiment.features({}) }
  end
  def test_reject_constant_feature_and_one_class
    constant=@training.map { |r| r.merge('attendance_pct'=>'50') }
    assert_raises(ArgumentError) { ExamExperiment.fit(constant) }
    one_class=@training.map { |r| r.merge('passed'=>'1') }
    assert_raises(ArgumentError) { ExamExperiment.fit(one_class) }
  end
  def test_bad_metric_inputs
    assert_raises(ArgumentError) { ExamExperiment.confusion([1],[0.5],1.1) }
    assert_raises(ArgumentError) { ExamExperiment.confusion([1],[]) }
  end
end
