aboutsummaryrefslogtreecommitdiff
path: root/test/test_classifier.rb
blob: a24c5ba1292cf4f1181e31e0808efbbde0d552fc (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
require 'linguist/classifier'
require 'linguist/language'
require 'linguist/samples'
require 'linguist/tokenizer'

require 'test/unit'

class TestClassifier < Test::Unit::TestCase
  include Linguist

  def samples_path
    File.expand_path("../../samples", __FILE__)
  end

  def fixture(name)
    File.read(File.join(samples_path, name))
  end

  def test_classify
    db = {}
    Classifier.train! db, "Ruby", fixture("Ruby/foo.rb")
    Classifier.train! db, "Objective-C", fixture("Objective-C/Foo.h")
    Classifier.train! db, "Objective-C", fixture("Objective-C/Foo.m")

    results = Classifier.classify(db, fixture("Objective-C/hello.m"))
    assert_equal "Objective-C", results.first[0]

    tokens  = Tokenizer.tokenize(fixture("Objective-C/hello.m"))
    results = Classifier.classify(db, tokens)
    assert_equal "Objective-C", results.first[0]
  end

  def test_restricted_classify
    db = {}
    Classifier.train! db, "Ruby", fixture("Ruby/foo.rb")
    Classifier.train! db, "Objective-C", fixture("Objective-C/Foo.h")
    Classifier.train! db, "Objective-C", fixture("Objective-C/Foo.m")

    results = Classifier.classify(db, fixture("Objective-C/hello.m"), ["Objective-C"])
    assert_equal "Objective-C", results.first[0]

    results = Classifier.classify(db, fixture("Objective-C/hello.m"), ["Ruby"])
    assert_equal "Ruby", results.first[0]
  end

  def test_instance_classify_empty
    results = Classifier.classify(Samples::DATA, "")
    assert results.first[1] < 0.5, results.first.inspect
  end

  def test_instance_classify_nil
    assert_equal [], Classifier.classify(Samples::DATA, nil)
  end

  def test_classify_ambiguous_languages
    Samples.each do |sample|
      language = Linguist::Language.find_by_name(sample[:language])
      next unless language.overrides.any?

      extname   = File.extname(sample[:path])
      languages = Language.all.select { |l| l.extensions.include?(extname) }.map(&:name)
      next unless languages.length > 1

      results = Classifier.classify(Samples::DATA, File.read(sample[:path]), languages)
      assert_equal language.name, results.first[0], "#{sample[:path]}\n#{results.inspect}"
    end
  end
end