Showing posts with label jgrapht. Show all posts
Showing posts with label jgrapht. Show all posts

Friday, October 24, 2008

Phrase Spelling Corrector using Word Collocation Probabilities

Spelling correction is one of those things that people don't notice when it works well. Indeed, for web-based search applications, its manifestation is usually a little "Did you mean: xxx?" component that appears when the application is not able to recognize the term being queried for. In spite of its relative non-ubiquity, however, users do notice when the suggestion is incorrect.

There are various approaches to spelling corrections. One popular approach is to use a Lucene index with character n-grams for terms in the index - I have written previously about my implementation of this approach.

Another popular approach is to compute edit costs and return from a dictionary the words that are within a predefined edit cost from the mispelt word. This is the approach used by GNU Aspell and its Java cousin Jazzy, which we use here.

Both these approaches work very well for single words, so they are very usable for applications such as word processors, where you need to be able to flag and suggest alternatives for mispelt words. In a typical search page, however, a user can type in a multi-word phrase, with one or more words mispelt. The job of the spelling corrector, in this case, is to tie the best suggestions together so that the corrected phrase makes sense within the context of the original phrase. A much harder problem, as you will no doubt agree.

Various approaches to solve this have been suggested and tried - I noticed one such suggestion almost by accident here on the Aspell TODO list, which set me thinking about this whole thing again.

Thinking about this suggestion a bit, I realized that a much simpler way would be to compute conditional probabilities between consecutive words in the phrase, and then consider the "best" suggestion to be the one which connects the words via the most probable path, i.e. the path with the highest sum of conditional probabilities. This effectively boils down a graph theory problem of computing the shortest path in a weighted directed graph. This post describes an implementation of this idea.

Consider Knud Sorensen's example from the Aspell TODO list. Two mispelt phrases and their corrected forms are shown below. As you can see, the correct form of the mispelt word 'fone' differs based on other words in the term.

1
2
    a fone number => a phone number
    a fone dress  => a fine dress

The list below shows the suggestions returned by Jazzy for the mispelt word 'fone', ordered by cost, i.e. the first suggestion is the one with the least edit cost to convert from the mispelt word. Notice that neither 'phone' nor 'fine' is the first result. The Java code for the CLI that I built for quickly looking up suggestions is available later in this post.

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
sujit@sirocco:~$ mvn -o exec:java \
  -Dexec.mainClass=com.mycompany.myapp.Shell \
  -Dexec.args=true
jazzy> fone
foe, one, fine, bone, zone, fore, lone, fond, font, hone, cone, gone, 
none, done, tone, fen, foes, on, fan, fin, fun, money, phone, son, fee, 
for, fog, fox, fined, found, fount, fence, fines, finer, honey, non, 
don, ton, ion, yon, vane, vine, June, gene, mane, mine, mono, bane, 
bony, pane, pine, pony, sane, sine, fade, fate, food, foot, vote, face, 
foci, fuse, fare, fire, four, free, fame, foam, fume, file, flee, foil, 
fool, foul, fowl, fake, lane, line, fife, five, fogs, find, fund, fans, 
fins, fang, cane, nine, dine, dune, tune, wane, wine
jazzy> \q

The approach I propose is to construct a graph of our input phrase ('a fone book'), adding vertices corresponding to the original word and each of its spelling suggestions, as shown below. The edge weights represent the conditional probability of the edge target B being followed by the edge source A (or P(B|A)). The numbers are all cooked up for this example, but I describe a way to compute them further down. I still need to populate my database tables with real data, I will describe this in a subsequent post.

What you will immediately notice is that we cannot prune the graph as we encounter each word in the phrase, i.e. we cannot select the most likely suggestion as we parse each word, since the "best path" is the most probable path through the graph from the start vertex to the finish vertex.

Since we are going to use Dijkstra's shortest path algorithm to find the shortest path (aka Graph Geodesic) through the graph, we need to convert the edge probabilities to a weight function given by wA,B, like so:

  wA,B = 1 - P(B|A)
  where:
    wA,B = cost to get from vertex A to B
    P(B|A) = probability of the occurrence of B given A

The probability P(B|A) can be computed as the number of times A and B co-occur in our dataset divided by the number of times word A occurs in the dataset, as shown below:

  If the occurrence of A and B are dependent:
    P(B ∩ A) = P(B|A) * P(A)
  so:
    P(B|A) = P(B ∩ A) / P(A)
           = N(B ∩ A) / N(A)

To get experimental values for N(A) and N(B ∩ A), we will need to extract data from actual search terms used by users from our Apache access logs and populate the following tables:

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
mysql> desc occur_a;
+---------+---------------+------+-----+---------+-------+
| Field   | Type          | Null | Key | Default | Extra |
+---------+---------------+------+-----+---------+-------+
| word    | varchar(32)   | NO   | PRI | NULL    |       | 
| n_words | mediumint(11) | NO   |     | NULL    |       | 
+---------+---------------+------+-----+---------+-------+

mysql> desc occur_ab;
+---------+---------------+------+-----+---------+-------+
| Field   | Type          | Null | Key | Default | Extra |
+---------+---------------+------+-----+---------+-------+
| word_a  | varchar(32)   | NO   | PRI | NULL    |       | 
| word_b  | varchar(32)   | NO   | PRI | NULL    |       | 
| n_words | mediumint(11) | NO   |     | NULL    |       | 
+---------+---------------+------+-----+---------+-------+

Without any data in the database tables, the code degrades very gracefully. It just returns what we typed in, as you can see below. This happens because we always insert the original word in the first position of the suggestion list returned by Jazzy, so the "best" among equals is the one that comes first. As before, the Java code for this CLI is provided later in the article.

1
2
3
4
5
6
sujit@sirocco:~$ mvn -o exec:java \
  -Dexec.mainClass=com.mycompany.myapp.Shell
spell-check> a fone book
a fone book
spell-check> a fone dress
a fone dress

Once some (still cooked up) occurrence data is entered manually into these tables...

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
mysql> select * from occur_a;
+-------+---------+
| word  | n_words |
+-------+---------+
| a     |     100 | 
| book  |      43 | 
| dress |      10 | 
| fine  |      12 | 
| phone |      18 | 
+-------+---------+
5 rows in set (0.00 sec)

mysql> select * from occur_ab;
+--------+--------+---------+
| word_a | word_b | n_words |
+--------+--------+---------+
| a      | fine   |       8 | 
| a      | phone  |      13 | 
| book   | phone  |      12 | 
| dress  | fine   |       7 | 
+--------+--------+---------+
4 rows in set (0.00 sec)

...our spelling corrector behaves more intelligently. The beauty of this approach is that its intelligence can be localized to your industry. So for example, if you were in the clothing business, your search terms are more likely to include fine dresses than phone books, and therefore the probability of P(dress|fine) would be higher than P(dress|phone).

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
sujit@sirocco:~$ mvn -o exec:java \
  -Dexec.mainClass=com.mycompany.myapp.Shell
spell-check> a fone book
a phone book
spell-check> a fone dress
a fine dress
spell-check> fone book
phone book
spell-check> fone dress
fine dress

Here is the code for the actual Spelling corrector. It uses Jazzy for its word suggestions, and JGraphT to construct a graph and run Dijkstra's shortest path algorithm (included in the JGraphT library) to find the most likely path based on word co-occurrence probabilities.

  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
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
// Source: src/main/java/com/mycompany/myapp/SpellingCorrector.java
package com.mycompany.myapp;

import java.io.File;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.List;

import javax.sql.DataSource;

import org.apache.commons.lang.StringUtils;
import org.jgrapht.alg.DijkstraShortestPath;
import org.jgrapht.graph.ClassBasedEdgeFactory;
import org.jgrapht.graph.DefaultWeightedEdge;
import org.jgrapht.graph.SimpleDirectedWeightedGraph;
import org.springframework.dao.IncorrectResultSizeDataAccessException;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.jdbc.datasource.DriverManagerDataSource;

import com.swabunga.spell.engine.SpellDictionary;
import com.swabunga.spell.engine.SpellDictionaryHashMap;
import com.swabunga.spell.engine.Word;

/**
 * Uses probability of word-collocations to determine best phrases to be
 * returned from a SpellingCorrector for multi-word mispelt queries.
 */
public class SpellingCorrector {

  private static final int SCORE_THRESHOLD = 200;
  private static final String DICTIONARY_FILENAME = 
    "src/main/resources/english.0";
  
  private long occurASumWords = 1L;
  private JdbcTemplate jdbcTemplate;
  
  @SuppressWarnings("unchecked")
  public String getSuggestion(String input) throws Exception {
    // initialize Jazzy spelling dictionary
    SpellDictionary dictionary = new SpellDictionaryHashMap(
      new File(DICTIONARY_FILENAME));
    // initialize database connection
    DataSource dataSource = new DriverManagerDataSource(
      "com.mysql.jdbc.Driver", "jdbc:mysql://localhost:3306/spelldb", 
      "foo", "secret");
    jdbcTemplate = new JdbcTemplate(dataSource);
    occurASumWords = jdbcTemplate.queryForLong(
      "select sum(n_words) from occur_a");
    if (occurASumWords == 0L) {
      // just a hack to prevent divide by zero for empty db
      occurASumWords = 1L;
    }
    // set up graph and create root vertex
    final SimpleDirectedWeightedGraph<SuggestedWord,DefaultWeightedEdge> g = 
      new SimpleDirectedWeightedGraph<SuggestedWord,DefaultWeightedEdge>(
      new ClassBasedEdgeFactory<SuggestedWord,DefaultWeightedEdge>(
      DefaultWeightedEdge.class));
    SuggestedWord startVertex = new SuggestedWord("START", 0);
    g.addVertex(startVertex);
    // set up variables to hold results of previous iteration
    List<SuggestedWord> prevVertices = 
      new ArrayList<SuggestedWord>();
    List<SuggestedWord> currentVertices = 
      new ArrayList<SuggestedWord>();
    int tokenId = 1;
    prevVertices.add(startVertex);
    // parse the string
    String[] tokens = input.toLowerCase().split("[ -]");
    for (String token : tokens) {
      // build up spelling suggestions for individual word
      List<String> possibleTokens = new ArrayList<String>();
      if (token.trim().length() <= 2) {
        // people usually don't make mistakes for words 2 words or less,
        // just pass it back unchanged
        possibleTokens.add(token);
      } else if (dictionary.isCorrect(token)) {
        // no need to find suggestions, token is recognized as valid spelling
        possibleTokens.add(token);
      } else {
        possibleTokens.add(token);
        List<Word> words = 
          dictionary.getSuggestions(token, SCORE_THRESHOLD);
        for (Word word : words) {
          possibleTokens.add(word.getWord());
        }
      }
      // populate the graph with these values
      for (String possibleToken : possibleTokens) {
        SuggestedWord currentVertex = 
          new SuggestedWord(possibleToken, tokenId); 
        g.addVertex(currentVertex);
        currentVertices.add(currentVertex);
        for (SuggestedWord prevVertex : prevVertices) {
          DefaultWeightedEdge edge = new DefaultWeightedEdge();
          double weight = computeEdgeWeight(
            prevVertex.token, currentVertex.token);
          g.setEdgeWeight(edge, weight);
          g.addEdge(prevVertex, currentVertex, edge);
        }
      }
      prevVertices.clear();
      prevVertices.addAll(currentVertices);
      currentVertices.clear();
      tokenId++;
    } // for token : tokens
    // finally set the end vertex
    SuggestedWord endVertex = new SuggestedWord("END", tokenId);
    g.addVertex(endVertex);
    for (SuggestedWord prevVertex : prevVertices) {
      DefaultWeightedEdge edge = new DefaultWeightedEdge();
      g.setEdgeWeight(edge, 1.0D);
      g.addEdge(prevVertex, endVertex, edge);
    }
    // find shortest path between START and END
    DijkstraShortestPath<SuggestedWord,DefaultWeightedEdge> dijkstra =
      new DijkstraShortestPath<SuggestedWord, DefaultWeightedEdge>(
      g, startVertex, endVertex);
    List<DefaultWeightedEdge> edges = dijkstra.getPathEdgeList();
    List<String> bestMatch = new ArrayList<String>();
    for (DefaultWeightedEdge edge : edges) {
      if (startVertex.equals(g.getEdgeSource(edge))) {
        // skip the START vertex
        continue;
      }
      bestMatch.add(g.getEdgeSource(edge).token);
    }
    return StringUtils.join(bestMatch.iterator(), " ");
  }

  private Double computeEdgeWeight(String prevToken, String currentToken) {
    if (prevToken.equals("START")) {
      // this is the first word, return 1-P(B)
      try {
        double nb = (Double) jdbcTemplate.queryForObject(
          "select n_words/? from occur_a where word = ?", 
          new Object[] {occurASumWords, currentToken}, Double.class);
        return 1.0D - nb;
      } catch (IncorrectResultSizeDataAccessException e) {
        // in case there is no match, then we should return weight of 1
        return 1.0D;
      }
    }
    double na = 0.0D;
    try {
      na = (Double) jdbcTemplate.queryForObject(
        "select n_words from occur_a where word = ?", 
        new String[] {prevToken}, Double.class);
    } catch (IncorrectResultSizeDataAccessException e) {
      // no match, should be 0
      na = 0.0D;
    }
    if (na == 0.0D) {
      // if N(A) == 0, A does not exist, and hence N(A ^ B) == 0 too,
      // so we guard against a DivideByZero and an additional useless
      // computation.
      return 1.0D;
    }
    // for the A^B lookup, alphabetize so A is lexically ahead of B
    // since that is the way we store it in the database
    String[] tokens = new String[] {prevToken, currentToken};
    Arrays.sort(tokens); // alphabetize before lookup
    double nba = 0.0D;
    try {
      nba = (Double) jdbcTemplate.queryForObject(
        "select n_words from occur_ab where word_a = ? and word_b = ?",
        tokens, Double.class);
    } catch (IncorrectResultSizeDataAccessException e) {
      // no result found so N(B^A) = 0
      nba = 0.0D;
    }
    return 1.0D - (nba / na);
  }

  /**
   * Holder for the graph vertex information.
   */
  private class SuggestedWord {
    public String token;
    public int id;
    
    public SuggestedWord(String token, int id) {
      this.token = token;
      this.id = id;
    }
    
    @Override
    public int hashCode() {
      return toString().hashCode();
    }
    
    @Override
    public boolean equals(Object obj) {
      if (!(obj instanceof SuggestedWord)) {
        return false;
      }
      SuggestedWord that = (SuggestedWord) obj;
      return (this.id == that.id && 
        this.token.equals(that.token));
    }
    
    @Override
    public String toString() {
      return id + ":" + token;
    }
  };
}

The CLI proved to be very useful for checking out assumptions quickly when I was developing the algorithm. Its quite simple, it just wraps the functionality within a JLine ConsoleReader. I included it here for completeness and to illustrate how easy it is to build. Depending on the presence of a command line argument, it can function either as an interface over the Jazzy dictionary or to the Phrase Spelling Corrector described in this post.

  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
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
// Source: src/main/java/com/mycompany/myapp/Shell.java
package com.mycompany.myapp;

import java.io.File;
import java.io.PrintWriter;
import java.util.ArrayList;
import java.util.Collections;
import java.util.Comparator;
import java.util.HashMap;
import java.util.List;
import java.util.Map;

import jline.ConsoleReader;

import org.apache.commons.lang.StringUtils;

import com.swabunga.spell.engine.SpellDictionary;
import com.swabunga.spell.engine.SpellDictionaryHashMap;
import com.swabunga.spell.engine.Word;

public class Shell {

  private final int SPELL_CHECK_THRESHOLD = 250;

  public Shell() throws Exception {
    ConsoleReader reader = new ConsoleReader(
      System.in, new PrintWriter(System.out));
    SpellingCorrector spellingCorrector = new SpellingCorrector();
    for (;;) {
      String line = reader.readLine("spell-check> ");
      if ("\\q".equals(line)) {
        break;
      }
      System.out.println(spellingCorrector.getSuggestion(line));
    }
  }

  // === this is really for exploratory testing purposes ===
  
  /**
   * Wrapper over Jazzy native spell checking functionality.
   * @param b always true (to differentiate from the new ctor).
   * @throws Exception if one is thrown.
   */
  public Shell(boolean b) throws Exception {
    ConsoleReader reader = new ConsoleReader(
      System.in, new PrintWriter(System.out));
    SpellDictionary dictionary = new SpellDictionaryHashMap(
      new File("src/main/resources/english.0"));
    for (;;) {
      String line = reader.readLine("jazzy> ");
      if ("\\q".equals(line)) {
        break;
      }
      String suggestions = suggest(dictionary, line);
      System.out.println(suggestions);
    }
  }
  
  /**
   * Looks up single words from Jazzy's English dictionary.
   * @param dictionary the dictionary object to look up.
   * @param incorrect the suspected mispelt word.
   * @return if the incorrect word is correct according to 
   * Jazzy's dictionary, then it is returned, else a set of possible
   * corrections is returned. If no possible corrections were found, 
   * this method returns (no suggestions).
   */
  @SuppressWarnings("unchecked")
  private String suggest(SpellDictionary dictionary, String incorrect) {
    if (dictionary.isCorrect(incorrect)) {
      // return the entered word
      return incorrect;
    }
    List<Word> words = dictionary.getSuggestions(
      incorrect, SPELL_CHECK_THRESHOLD);
    List<String> suggestions = new ArrayList<String>();
    final Map<String,Integer> costs = 
      new HashMap<String,Integer>();
    for (Word word : words) {
      costs.put(word.getWord(), word.getCost());
      suggestions.add(word.getWord());
    }
    if (suggestions.size() == 0) {
      return "(no suggestions)";
    }
    Collections.sort(suggestions, new Comparator<String>() {
      public int compare(String s1, String s2) {
        Integer cost1 = costs.get(s1);
        Integer cost2 = costs.get(s2);
        return cost1.compareTo(cost2);
      }
    });
    return StringUtils.join(suggestions.iterator(), ", ");
  }
  
  public static void main(String[] args) throws Exception {
    if (args.length == 1) {
      // word mode
      new Shell(true);
    } else {
      new Shell();
    }
  }
}

Note that I still don't know whether this works well for a large set of mispelt phrases. I need to put this through a lot more real data to say that with any degree of certainty. It is also fairly slow in my development testing. I have a few ideas as to how that can be improved, although I will attempt them after I have some real data to play with. As always, any suggestions/corrections much appreciated.

Update 2009-04-26: In recent posts, I have been building on code written and described in previous posts, so there were (and rightly so) quite a few requests for the code. So I've created a project on Sourceforge to host the code. You will find the complete source code built so far in the project's SVN repository.

Saturday, May 31, 2008

Modeling an Ontology in memory with JGraphT

Last week, I blogged about a custom StAX parser that parsed an OWL XML file representing an ontology of wine. The parser parsed out the information into a MySQL database. The database structure has changed slightly, with a couple of tables being renamed to make the design a bit more obvious. The new schema looks like this:

What this database is trying to model is a bunch of facts modeled as semantic triples. A semantic triple consists of two entities connected by a relationship. The entity object only has an id and name and a list of Attribute objects. This is so we can beef up our Entity over time, as we discover more properties for these objects, without having to change any code. Attributes are modeled as name-value tuples, and the AttributeType normalizes the attribute names, which are likely to be repeated across Entities.

One thing we did before we go forward is to add reverse relationships. After the OWL file was parsed and loaded into the database, we ended up with about 9 relationship types. These represent one way relationships (such as subClassOf). Usually relationships are two way, so we manually set the reverse relationships with the id as the negative of the original relationId. The complete list of relationships is shown below. Of course, we need not put in relationships that don't make sense or that we don't want to expose.

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
+----+---------------------+
| id | name                |
+----+---------------------+
| -9 | colorProperty       | 
| -8 | vintageYearProperty | 
| -7 | mainIngredient      | 
| -6 | bodyProperty        | 
| -5 | flavorProperty      | 
| -4 | sugarProperty       | 
| -3 | makes               | 
| -2 | contains            | 
| -1 | superclassOf        | 
|  1 | subclassOf          | 
|  2 | locatedIn           | 
|  3 | hasMaker            | 
|  4 | hasSugar            | 
|  5 | hasFlavor           | 
|  6 | hasBody             | 
|  7 | madeFromGrape       | 
|  8 | hasVintageYear      | 
|  9 | hasColor            | 
+----+---------------------+

An ontology can be visualized as a forest of taxonomy trees, where the nodes of the trees are connected to nodes of other trees - in other words, a graph. So my next step is to convert this structure into an in-memory graph object so it can be navigated without having to resort to complex SQL.

Searching for decent Java based graph data structures I could use, I came upon JGraphT, which not only provides standard graph data structures that can be used, but also has a large number of graph algorithms built into the package. I guess I could have cooked one up myself, since all I wanted to do was to model a graph and navigate it, but the advantage of using a standard data structure from a decent library is that the library author has already worked out the kinks in the data structure so it is likely to be more extensible. Moreover, while I don't need any of the graph algorithms built into JGraphT right now, it is conceivable that I will at some point down the road.

So anyway, this post describes the code that I wrote to load a JGraphT Graph object from my database, and then hitting the graph with a few basic queries to make sure everything works.

First the beans. I define an Entity bean, an Attribute bean, a Relation bean, and a Fact bean which models a semantic triple. These are simple holder classes.

 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
69
70
71
72
73
74
75
// Entity.java
package com.mycompany.myapp.ontology;

import java.io.Serializable;
import java.util.ArrayList;
import java.util.List;

import org.apache.commons.lang.builder.EqualsBuilder;
import org.apache.commons.lang.builder.ReflectionToStringBuilder;
import org.apache.commons.lang.builder.ToStringStyle;

public class Entity implements Serializable {
  
  private static final long serialVersionUID = 54272228896206677L;

  private long id;
  private String name;
  private List<Attribute> attributes = new ArrayList<Attribute>();
  
  public Entity() {
    super();
  }
  
  public Entity(long id) {
    this();
    setId(id);
  }
  
  public long getId() {
    return id;
  }

  public void setId(long id) {
    this.id = id;
  }
  
  public String getName() {
    return name;
  }
  
  public void setName(String name) {
    this.name = name;
  }
  
  public List<Attribute> getAttributes() {
    return attributes;
  }

  public void setAttributes(List<Attribute> attributes) {
    this.attributes = attributes;
  }

  public void addAttribute(Attribute attribute) {
    this.attributes.add(attribute);
  }

  @Override
  public int hashCode() {
    return (int) id;
  }
  
  @Override
  public boolean equals(Object obj) {
    if (!(obj instanceof Entity)) {
      return false;
    }
    Entity that = (Entity) obj;
    return EqualsBuilder.reflectionEquals(this, that);
  }
  
  @Override
  public String toString() {
    return ReflectionToStringBuilder.reflectionToString(this, ToStringStyle.NO_FIELD_NAMES_STYLE);
  }
}
 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
// Attribute.java
package com.mycompany.myapp.ontology;

public class Attribute {
  
  private String name;
  private String value;

  public Attribute() {
    super();
  }
  
  public Attribute(String name, String value) {
    this();
    setName(name);
    setValue(value);
  }
  
  public String getName() {
    return name;
  }
  
  public void setName(String name) {
    this.name = name;
  }
  
  public String getValue() {
    return value;
  }
  
  public void setValue(String value) {
    this.value = value;
  }
}
 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
// Relation.java
package com.mycompany.myapp.ontology;

import java.io.Serializable;

import org.apache.commons.lang.builder.ReflectionToStringBuilder;
import org.apache.commons.lang.builder.ToStringStyle;

public class Relation implements Serializable {

  private static final long serialVersionUID = 8521110824988338681L;
  
  private long relationId;
  private String relationName;
  
  public Relation() {
    super();
  }

  public long getId() {
    return relationId;
  }

  public void setRelationId(long relationId) {
    this.relationId = relationId;
  }

  public String getName() {
    return relationName;
  }

  public void setRelationName(String relationName) {
    this.relationName = relationName;
  }

  @Override
  public String toString() {
    return ReflectionToStringBuilder.reflectionToString(this, ToStringStyle.NO_FIELD_NAMES_STYLE);
  }
}
 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
// Fact.java
package com.mycompany.myapp.ontology;

public class Fact {

  private long sourceEntityId;
  private long targetEntityId;
  private long relationId;
  
  public Fact() {
    super();
  }
  
  public Fact(long sourceEntityId, long targetEntityId, long relationId) {
    this();
    setSourceEntityId(sourceEntityId);
    setTargetEntityId(targetEntityId);
    setRelationId(relationId);
  }
  
  public long getSourceEntityId() {
    return sourceEntityId;
  }
  
  public void setSourceEntityId(long sourceEntityId) {
    this.sourceEntityId = sourceEntityId;
  }
  
  public long getTargetEntityId() {
    return targetEntityId;
  }
  
  public void setTargetEntityId(long targetEntityId) {
    this.targetEntityId = targetEntityId;
  }
  
  public long getRelationId() {
    return relationId;
  }

  public void setRelationId(long relationId) {
    this.relationId = relationId;
  }
}

The overriden equals() and hashCode() methods on the Entity bean (above) are necessary - this is to enable JGraphT to locate it in the graph when we try to look for it with a reference to an Entity object.

To connect Entities, we need an Edge object that can be labelled, so we subclass JGraphT's DefaultEdge and add in an additional property relationId. Here is the code for RelationEdge.java.

 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
// RelationEdge.java
package com.mycompany.myapp.ontology;

import org.apache.commons.lang.builder.EqualsBuilder;
import org.apache.commons.lang.builder.ReflectionToStringBuilder;
import org.apache.commons.lang.builder.ToStringStyle;
import org.jgrapht.graph.DefaultEdge;

/**
 * Extends DefaultEdge to add a label to the graph. The label is the
 * relationId that relates the entities the edge connects.
 */
public class RelationEdge extends DefaultEdge {
  
  private static final long serialVersionUID = 1994877217677659613L;

  private long relationId;

  public RelationEdge() {
    super();
  }
  
  public RelationEdge(long relationId) {
    this();
    setRelationId(relationId);
  }
  
  public long getRelationId() {
    return relationId;
  }

  public void setRelationId(long relationId) {
    this.relationId = relationId;
  }
}

Finally, we define the container class which ties this all together. The client will call methods on this class.

  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
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
// Ontology.java
package com.mycompany.myapp.ontology;

import java.io.Serializable;
import java.util.HashMap;
import java.util.HashSet;
import java.util.Map;
import java.util.Set;

import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.jgrapht.Graph;
import org.jgrapht.graph.ClassBasedEdgeFactory;
import org.jgrapht.graph.SimpleDirectedGraph;

public class Ontology implements Serializable {

  private static final long serialVersionUID = 8903265933795172508L;
  
  private final Log log = LogFactory.getLog(getClass());
  
  protected Map<Long,Entity> entityMap;
  protected Map<Long,Relation> relationMap;
  protected SimpleDirectedGraph<Entity,RelationEdge> ontology;

  public Ontology() {
    entityMap = new HashMap<Long,Entity>();
    relationMap = new HashMap<Long,Relation>();
    ontology = new SimpleDirectedGraph<Entity,RelationEdge>(
      new ClassBasedEdgeFactory<Entity,RelationEdge>(RelationEdge.class));
  }

  public Entity getEntityById(long entityId) {
    return entityMap.get(entityId);
  }

  public Relation getRelationById(long relationId) {
    return relationMap.get(relationId);
  }
  
  public Set<Long> getAvailableRelationIds(Entity entity) {
    Set<Long> relationIds = new HashSet<Long>();
    Set<RelationEdge> relationEdges = ontology.edgesOf(entity);
    for (RelationEdge relationEdge : relationEdges) {
      relationIds.add(relationEdge.getRelationId());
    }
    return relationIds;
  }
  
  public Set<Entity> getEntitiesRelatedById(Entity entity, long relationId) {
    Set<RelationEdge> relationEdges = ontology.outgoingEdgesOf(entity);
    Set<Entity> relatedEntities = new HashSet<Entity>();
    for (RelationEdge relationEdge : relationEdges) {
      if (relationEdge.getRelationId() == relationId) {
        Entity relatedEntity = ontology.getEdgeTarget(relationEdge);
        relatedEntities.add(relatedEntity);
      }
    }
    return relatedEntities;
  }
  
  public void addEntity(Entity entity) {
    entityMap.put(entity.getId(), entity);
    ontology.addVertex(entity);
  }
  
  public void addRelation(Relation relation) throws Exception {
    relationMap.put(relation.getId(), relation);
  }
  
  public void addFact(Fact fact) throws Exception {
    Entity sourceEntity = getEntityById(fact.getSourceEntityId());
    if (sourceEntity == null) {
      log.error("No entity found for source entityId:" + fact.getSourceEntityId());
      return;
    }
    Entity targetEntity = getEntityById(fact.getTargetEntityId());
    if (targetEntity == null) {
      log.error("No entity found for target entityId: " + fact.getTargetEntityId());
      return;
    }
    long relationId = fact.getRelationId();
    Relation relation = getRelationById(relationId);
    if (relation == null) {
      log.error("No relation found for relationId: " + relationId);
      return;
    }
    // does fact exist? If so, dont do anything, just return
    Set<Long> relationIds = getAvailableRelationIds(sourceEntity);
    if (relationIds.contains(relationId)) {
      log.info("Fact: " + relation.getName() + "(" + 
        sourceEntity.getName() + "," + targetEntity.getName() + 
        ") already added to ontology");
      return;
    }
    RelationEdge relationEdge = new RelationEdge();
    relationEdge.setRelationId(relationId);
    ontology.addEdge(sourceEntity, targetEntity, relationEdge);
    if (relationMap.get(-1L * relationId) != null) {
      RelationEdge reverseRelationEdge = new RelationEdge();
      reverseRelationEdge.setRelationId(-1L * relationId);
      ontology.addEdge(targetEntity, sourceEntity, reverseRelationEdge);
    }
  }
}

To load this object, we use the DbOntologyLoader class, which calls methods on the DAO classes to retrieve data from the database. Here is the code for the loader.

 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
// DbOntologyLoader.java
package com.mycompany.myapp.ontology.loaders;

import com.mycompany.myapp.ontology.daos.EntityDao;
import com.mycompany.myapp.ontology.daos.FactDao;
import com.mycompany.myapp.ontology.daos.RelationDao;
import com.mycompany.myapp.ontology.Fact;
import com.mycompany.myapp.ontology.Ontology;
import com.mycompany.myapp.ontology.Entity;
import com.mycompany.myapp.ontology.Relation;

import java.util.List;

import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;

public class DbOntologyLoader {
  
  private final Log log = LogFactory.getLog(getClass());
  
  private EntityDao entityDao;
  private RelationDao relationDao;
  private FactDao factDao;
  
  public void setEntityDao(EntityDao entityDao) {
    this.entityDao = entityDao;
  }

  public void setRelationDao(RelationDao relationDao) {
    this.relationDao = relationDao;
  }

  public void setFactDao(FactDao factDao) {
    this.factDao = factDao;
  }

  public Ontology load() throws Exception {
    Ontology ontology = new Ontology();
    log.debug("Loading entities");
    List<Entity> entities = entityDao.getAllEntities();
    for (Entity entity : entities) {
      ontology.addEntity(entity);
    }
    log.debug("Loading relations");
    List<Relation> relations = relationDao.getAllRelations();
    for (Relation relation : relations) {
      ontology.addRelation(relation);
      if (relationDao.isBidirectional(relation.getId())) {
        Relation reverseRelation = relationDao.getById(-1L * relation.getId());
        ontology.addRelation(reverseRelation);
      }
    }
    log.debug("Loading facts");
    List<Fact> facts = factDao.getAllFacts();
    for (Fact fact : facts) {
      ontology.addFact(fact);
    }
    log.debug("Ontology load complete");
    return ontology;
  }
}

The loader depends on three DAOs for Entity, Relation and Fact. The code for these is shown below for completeness.

  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
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
// EntityDao.java
package com.mycompany.myapp.ontology.daos;

import java.sql.Connection;
import java.sql.PreparedStatement;
import java.sql.SQLException;
import java.sql.Statement;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;

import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.springframework.dao.IncorrectResultSizeDataAccessException;
import org.springframework.jdbc.core.PreparedStatementCreator;
import org.springframework.jdbc.core.support.JdbcDaoSupport;
import org.springframework.jdbc.support.GeneratedKeyHolder;
import org.springframework.jdbc.support.KeyHolder;

import com.mycompany.myapp.ontology.Attribute;
import com.mycompany.myapp.ontology.Entity;

public class EntityDao extends JdbcDaoSupport {

  private final Log log = LogFactory.getLog(getClass());

  @SuppressWarnings("unchecked")
  public List<Entity> getAllEntities() {
    List<Entity> entities = new ArrayList<Entity>();
    List<Map<String,Object>> rows = getJdbcTemplate().queryForList(
      "select id, name from entities");
    for (Map<String,Object> row : rows) {
      Entity entity = new Entity();
      entity.setId((Integer) row.get("ID"));
      entity.setName((String) row.get("NAME"));
      entities.add(entity);
    }
    return entities;
  }

  @SuppressWarnings("unchecked")
  public Entity getById(long id) {
    try {
      Entity entity = new Entity();
      Map<String,Object> row = getJdbcTemplate().queryForMap(
        "select id, name from entities where id = ?", 
        new Long[] {id});
      entity.setId((Integer) row.get("ID"));
      entity.setName((String) row.get("NAME"));
      return entity;
    } catch (IncorrectResultSizeDataAccessException e) {
      return null;
    }
  }
  
  @SuppressWarnings("unchecked")
  public Entity getByName(String name) {
    try {
      Entity entity = new Entity();
      Map<String,Object> row = getJdbcTemplate().queryForMap(
        "select id, name from entities where name = ?", 
        new String[] {name});
      entity.setId((Integer) row.get("ID"));
      entity.setName((String) row.get("NAME"));
      entity.setAttributes(getAttributes(entity.getId()));
      return entity;
    } catch (IncorrectResultSizeDataAccessException e) {
      return null;
    }
  }
  
  @SuppressWarnings("unchecked")
  public List<Attribute> getAttributes(long entityId) {
    List<Attribute> attributes = new ArrayList<Attribute>();
    List<Map<String,String>> rows = getJdbcTemplate().queryForList(
      "select at.attr_name, a.value " +
      "from attributes a, attribute_types at " +
      "where a.attr_id = at.id " +
      "and a.entity_id = ?", new Long[] {entityId});
    for (Map<String,String> row : rows) {
      String name = row.get("ATTR_NAME");
      String value = row.get("VALUE");
      Attribute attribute = new Attribute(name, value);
      attributes.add(attribute);
    }
    return attributes;
  }

  @SuppressWarnings("unchecked")
  public Attribute getAttributeByName(long entityId, String attributeName) {
    try {
      Attribute attribute = new Attribute();
      Map<String,String> row = getJdbcTemplate().queryForMap(
        "select at.attr_name, a.value " +
        "from attributes a, attribute_types at " +
        "where a.attr_id = at.id " +
        "and a.entity_id = ? " +
        "and at.attr_name = ?", new Object[] {entityId, attributeName});
      attribute.setName(row.get("NAME"));
      attribute.setValue(row.get("VALUE"));
      return attribute;
    } catch (IncorrectResultSizeDataAccessException e) {
      return null;
    }
  }

  public long getAttributeTypeId(final String attributeName) {
    try {
      long attributeTypeId = getJdbcTemplate().queryForLong(
        "select id from attribute_types where attr_name = ?", 
        new String[] {attributeName});
      return attributeTypeId;
    } catch (IncorrectResultSizeDataAccessException e) {
      return 0L;
    }
  }
  
  public long save(final Entity entity) {
    Entity dbEntity = getByName(entity.getName());
    if (dbEntity == null) {
      log.debug("Saving entity:" + entity.getName());
      // insert the entity
      KeyHolder entityKeyHolder = new GeneratedKeyHolder();
      getJdbcTemplate().update(new PreparedStatementCreator() {
        public PreparedStatement createPreparedStatement(Connection conn)
        throws SQLException {
          PreparedStatement ps = conn.prepareStatement(
            "insert into entities(name) values (?)", 
            Statement.RETURN_GENERATED_KEYS);
          ps.setString(1, entity.getName());
          return ps;
        }
      }, entityKeyHolder);
      long entityId = entityKeyHolder.getKey().longValue();
      List<Attribute> attributes = entity.getAttributes();
      for (Attribute attribute : attributes) {
        saveAttribute(entityId, attribute);
      }
      // finally, always save the "english name" of the entity as an attribute
      saveAttribute(entityId, new Attribute("EnglishName", getEnglishName(entity)));
      return entityId;
    } else {
      getJdbcTemplate().update("update entities set name = ? where id = ?", 
        new Object[] {entity.getName(), entity.getId()});
      return entity.getId();
    }
  }

  public long saveAttribute(final long entityId, final Attribute attribute) {
    // check to see if attribute exists in attribute_types
    long attributeTypeId = getAttributeTypeId(attribute.getName());
    if (attributeTypeId == 0L) {
      attributeTypeId = saveAttributeType(attribute.getName());
    }
    Attribute dbAttribute = getAttributeByName(entityId, attribute.getName());
    final long attrId = attributeTypeId;
    if (dbAttribute == null) {
      KeyHolder keyholder = new GeneratedKeyHolder();
      final String attributeName = attribute.getName();
      getJdbcTemplate().update(new PreparedStatementCreator() {
        public PreparedStatement createPreparedStatement(Connection conn)
        throws SQLException {
          PreparedStatement ps = conn.prepareStatement(
            "insert into attributes(entity_id, attr_id, value) values (?, ?, ?)");
          ps.setLong(1, entityId);
          ps.setLong(2, attrId);
          ps.setString(3, attribute.getValue());
          return ps;
        }
      }, keyholder);
      long attributeId = keyholder.getKey().longValue();
      return attributeId;
    } else {
      getJdbcTemplate().update(
        "update attributes set value = ? where entity_id = ? and attr_id = ?", 
        new Long[] {entityId, attrId});
      return attrId;
    }
  }

  public long saveAttributeType(final String attributeName) {
    long attributeTypeId = getAttributeTypeId(attributeName);
    if (attributeTypeId == 0L) {
      KeyHolder keyholder = new GeneratedKeyHolder();
      getJdbcTemplate().update(new PreparedStatementCreator() {
        public PreparedStatement createPreparedStatement(Connection conn)
        throws SQLException {
          PreparedStatement ps = conn.prepareStatement(
            "insert into attribute_types(attr_name) values (?)");
          ps.setString(1, attributeName);
          return ps;
        }
      }, keyholder);
      attributeTypeId = keyholder.getKey().longValue();
    }
    return attributeTypeId;
  }
    
  /**
   * Split up Uppercase Camelcased names (like Java classnames or C++ variable
   * names) into English phrases by splitting wherever there is a transition 
   * from lowercase to uppercase.
   * @param name the input camel cased name.
   * @return the "english" name.
   */
  public String getEnglishName(Entity entity) {
    if (entity == null) {
      return null;
    }
    StringBuilder englishNameBuilder = new StringBuilder();
    char[] namechars = entity.getName().toCharArray();
    for (int i = 0; i < namechars.length; i++) {
      if (i > 0 && Character.isUpperCase(namechars[i]) && 
          Character.isLowerCase(namechars[i-1])) {
        englishNameBuilder.append(' ');
      }
      englishNameBuilder.append(namechars[i]);
    }
    return englishNameBuilder.toString();
  }
}
 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
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
// RelationDao.java
package com.mycompany.myapp.ontology.daos;

import java.sql.Connection;
import java.sql.PreparedStatement;
import java.sql.SQLException;
import java.sql.Statement;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;

import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.springframework.dao.IncorrectResultSizeDataAccessException;
import org.springframework.jdbc.core.PreparedStatementCreator;
import org.springframework.jdbc.core.support.JdbcDaoSupport;
import org.springframework.jdbc.support.GeneratedKeyHolder;
import org.springframework.jdbc.support.KeyHolder;

import com.mycompany.myapp.ontology.Relation;

public class RelationDao extends JdbcDaoSupport {

  private final Log log = LogFactory.getLog(getClass());
  
  @SuppressWarnings("unchecked")
  public List<Relation> getAllRelations() {
    List<Relation> relations = new ArrayList<Relation>();
    List<Map<String,Object>> rows = getJdbcTemplate().queryForList(
      "select id, name from relations where id > 0");
    for (Map<String,Object> row : rows) {
      Relation relation = new Relation();
      relation.setRelationId((Integer) row.get("ID"));
      relation.setRelationName((String) row.get("NAME"));
      relations.add(relation);
    }
    return relations;
  }

  @SuppressWarnings("unchecked")
  public Relation getById(long relationId) {
    Relation relation = new Relation();
    try {
      Map<String,Object> row = getJdbcTemplate().queryForMap(
        "select id, name from relations where id = ?", 
        new Long[] {relationId});
      relation.setRelationId((Integer) row.get("ID"));
      relation.setRelationName((String) row.get("NAME"));
      return relation;
    } catch (IncorrectResultSizeDataAccessException e) {
      return null;
    }
  }

  @SuppressWarnings("unchecked")
  public Relation getByName(String name) {
    Relation relation = new Relation();
    try {
      Map<String,Object> row = getJdbcTemplate().queryForMap(
        "select id, name from relations where id = ?", 
        new String[] {name});
      relation.setRelationId((Integer) row.get("ID"));
      relation.setRelationName((String) row.get("NAME"));
      return relation;
    } catch (IncorrectResultSizeDataAccessException e) {
      return null;
    }
  }
  
  public boolean isBidirectional(long relationId) {
    int count = getJdbcTemplate().queryForInt(
      "select count(*) from relations where id = ?", 
      new Long[] {-1L * relationId});
    return count > 0;
  }
  
  public long save(final Relation relation) {
    Relation dbRelation = getByName(relation.getName());
    if (dbRelation == null) {
      KeyHolder keyholder = new GeneratedKeyHolder();
      getJdbcTemplate().update(new PreparedStatementCreator() {
        public PreparedStatement createPreparedStatement(Connection conn) 
            throws SQLException {
          PreparedStatement ps = conn.prepareStatement(
            "insert into relations(name) values (?)", 
            Statement.RETURN_GENERATED_KEYS);
          ps.setString(1, relation.getName());
          return ps;
        }
      }, keyholder);
      return keyholder.getKey().longValue();
    } else {
      getJdbcTemplate().update("update relations set name = ? where id = ?",
        new Object[] {relation.getName(), relation.getId()});
      return relation.getId();
    }
  }
}
  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
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
// FactDao.java
package com.mycompany.myapp.ontology.daos;

import java.sql.Connection;
import java.sql.PreparedStatement;
import java.sql.SQLException;
import java.sql.Statement;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;

import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.springframework.dao.IncorrectResultSizeDataAccessException;
import org.springframework.jdbc.core.PreparedStatementCreator;
import org.springframework.jdbc.core.support.JdbcDaoSupport;
import org.springframework.jdbc.support.GeneratedKeyHolder;
import org.springframework.jdbc.support.KeyHolder;

import com.mycompany.myapp.ontology.Entity;
import com.mycompany.myapp.ontology.Fact;
import com.mycompany.myapp.ontology.Relation;

public class FactDao extends JdbcDaoSupport {

  private final Log log = LogFactory.getLog(getClass());
  
  private EntityDao entityDao;
  private RelationDao relationDao;
  
  public void setEntityDao(EntityDao entityDao) {
    this.entityDao = entityDao;
  }
  
  public void setRelationDao(RelationDao relationDao) {
    this.relationDao = relationDao;
  }

  @SuppressWarnings("unchecked")
  public List<Fact> getAllFacts() {
    List<Fact> facts = new ArrayList<Fact>();
    List<Map<String,Integer>> rows = getJdbcTemplate().queryForList(
      "select f.src_entity_id, f.trg_entity_id, f.relation_id " +
      "from facts f, relations r " +
      "where f.relation_id = r.id");
    for (Map<String,Integer> row : rows) {
      Fact fact = new Fact();
      fact.setSourceEntityId(row.get("SRC_ENTITY_ID"));
      fact.setTargetEntityId(row.get("TRG_ENTITY_ID"));
      fact.setRelationId(row.get("RELATION_ID"));
      facts.add(fact);
    }
    return facts;
  }

  public void save(Fact fact) {
    Entity sourceEntity = entityDao.getById(fact.getSourceEntityId());
    Entity targetEntity = entityDao.getById(fact.getTargetEntityId());
    if (sourceEntity == null || targetEntity == null) {
      log.error("Cannot relate null entities");
      return;
    }
    Relation relation = relationDao.getById(fact.getRelationId());
    if (relation == null) {
      log.error("Unknown relation, cannot save fact");
      return;
    }
    save(sourceEntity.getName(), targetEntity.getName(), relation.getName());
  }
  
  public void save(final String sourceEntityName, final String targetEntityName, 
      final String relationName) {
    // get the entity ids for source and target
    Entity sourceEntity = entityDao.getByName(sourceEntityName);
    Entity targetEntity = entityDao.getByName(targetEntityName);
    if (sourceEntity == null || targetEntity == null) {
      log.error("Cannot save relation: " + relationName + "(" + 
        sourceEntityName + "," + targetEntityName + ")"); 
      return;
    }
    log.debug("Saving relation: " + relationName + "(" + 
      sourceEntityName + "," + targetEntityName + ")");
    // get the relation id
    long relationTypeId = 0L;
    try {
      relationTypeId = getJdbcTemplate().queryForInt(
        "select id from relations where name = ?", 
        new String[] {relationName});
    } catch (IncorrectResultSizeDataAccessException e) {
      KeyHolder keyholder = new GeneratedKeyHolder();
      getJdbcTemplate().update(new PreparedStatementCreator() {
        public PreparedStatement createPreparedStatement(Connection conn) 
            throws SQLException {
          PreparedStatement ps = conn.prepareStatement(
            "insert into relations(name) values (?)", 
            Statement.RETURN_GENERATED_KEYS);
          ps.setString(1, relationName);
          return ps;
        }
      }, keyholder);
      relationTypeId = keyholder.getKey().longValue();
    }
    // save it
    getJdbcTemplate().update(
      "insert into facts(src_entity_id, trg_entity_id, relation_id) values (?, ?, ?)", 
      new Long[] {sourceEntity.getId(), targetEntity.getId(), relationTypeId});
  }
}

Finally, to test this thing out, we start up the loader and populate our graph, then issue queries against it (in the form of JUnit test methods) and see what we get. Here is the JUnit test:

 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
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
// OntologyTest.java
package com.mycompany.myapp.ontology;

import java.util.Set;

import javax.sql.DataSource;

import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.jgrapht.Graph;
import org.junit.Assert;
import org.junit.BeforeClass;
import org.junit.Test;
import org.springframework.jdbc.datasource.DriverManagerDataSource;

import com.mycompany.myapp.ontology.daos.EntityDao;
import com.mycompany.myapp.ontology.daos.FactDao;
import com.mycompany.myapp.ontology.daos.RelationDao;
import com.mycompany.myapp.ontology.loaders.DbOntologyLoader;

public class OntologyTest {

  private final Log log = LogFactory.getLog(getClass());
  
  private static Ontology ontology;
  
  @BeforeClass
  public static void setUpBeforeClass() throws Exception {
    
    DataSource dataSource = new DriverManagerDataSource(
        "com.mysql.jdbc.Driver", "jdbc:mysql://localhost:3306/ontodb", "ontodev", "****");
    
    EntityDao entityDao = new EntityDao();
    entityDao.setDataSource(dataSource);
    
    RelationDao relationDao = new RelationDao();
    relationDao.setDataSource(dataSource);
    
    FactDao factDao = new FactDao();
    factDao.setDataSource(dataSource);
    factDao.setEntityDao(entityDao);
    factDao.setRelationDao(relationDao);
    
    DbOntologyLoader loader = new DbOntologyLoader();
    loader.setEntityDao(entityDao);
    loader.setRelationDao(relationDao);
    loader.setFactDao(factDao);
    
    ontology = loader.load();
  }
  
  @Test
  public void testLoad() throws Exception {
    Graph<Entity,RelationEdge> ontologyGraph = ontology.ontology;
    // We should have 237 vertices and about 500 edges
    log.debug("# vertices =" + ontologyGraph.vertexSet().size());
    Assert.assertTrue("#-vertices test failed", ontologyGraph.vertexSet().size() == 237);
    log.debug("# edges = " + ontologyGraph.edgeSet().size());
    Assert.assertTrue("#-edges test failed", ontologyGraph.edgeSet().size() == 500);
  }
  
  @Test
  public void testWhereIsLoireRegion() throws Exception {
    Entity loireRegion = ontology.getEntityById(26);
    long locatedInRelationId = 2L;
    Set<Entity> entities = ontology.getEntitiesRelatedById(loireRegion, locatedInRelationId);
    log.debug("query> where is Loire Region?");
    for (Entity entity : entities) {
      log.debug("..." + entity.getName());
    }
  }
  
  @Test
  public void testWhatRegionsAreInUSRegion() throws Exception {
    Entity usRegion = ontology.getEntityById(23);
    long reverseLocatedInRelationId = -2L;
    Set<Entity> entities = ontology.getEntitiesRelatedById(usRegion, reverseLocatedInRelationId);
    log.debug("query> what regions are in US Region?");
    for (Entity entity : entities) {
      log.debug("..." + entity.getName());
    }
  }
  
  @Test
  public void testWhatAreSweetWines() throws Exception {
    Entity sweetWinesEntity = ontology.getEntityById(125);
    long reverseOfHasSugarRelationId = -4L;
    Set<Entity> entities = ontology.getEntitiesRelatedById(
      sweetWinesEntity, reverseOfHasSugarRelationId);
    log.debug("query> what are sweet wines?");
    for (Entity entity : entities) {
      log.debug("..." + entity.getName());
    }
  }
}

And here are the actual outputs (formatted for clarity) for the questions we (effectively) asked in our unit tests above, and the answers we got back from our ontology.

1
2
3
4
5
6
7
8
9
 query> where is Loire Region?
 ...FrenchRegion
 query> what regions are in US Region?
 ...CaliforniaRegion
 ...TexasRegion
 query> what are sweet wines?
 ...WhitehallLanePrimavera
 ...SchlossVolradTrochenbierenausleseRiesling
 ...SchlossRothermelTrochenbierenausleseRiesling

As my son would say -- "Cool, huh?"

Update 2009-04-26: In recent posts, I have been building on code written and described in previous posts, so there were (and rightly so) quite a few requests for the code. So I've created a project on Sourceforge to host the code. You will find the complete source code built so far in the project's SVN repository.