diff --git a/sqlTranslate/.idea/workspace.xml b/sqlTranslate/.idea/workspace.xml
index 66f4221ad..c3aaaf2ce 100644
--- a/sqlTranslate/.idea/workspace.xml
+++ b/sqlTranslate/.idea/workspace.xml
@@ -4,14 +4,12 @@
-
-
-
-
-
-
-
+
+
+
+
+
@@ -410,7 +408,14 @@
1726148093221
-
+
+ 1726149015586
+
+
+
+ 1726149015586
+
+
@@ -448,7 +453,8 @@
-
+
+
diff --git a/sqlTranslate/src/main/java/Generator/OpenGaussGenerator.java b/sqlTranslate/src/main/java/Generator/OpenGaussGenerator.java
index 87cd440cf..c5acf0dee 100644
--- a/sqlTranslate/src/main/java/Generator/OpenGaussGenerator.java
+++ b/sqlTranslate/src/main/java/Generator/OpenGaussGenerator.java
@@ -12,6 +12,8 @@ import Parser.AST.Join.JoinConditionNode;
import Parser.AST.Join.JoinSourceTabNode;
import Parser.AST.Select.SelectNode;
import Parser.AST.Select.SelectObjNode;
+import Parser.AST.Update.UpdateNode;
+import Parser.AST.Update.UpdateObjNode;
import Parser.OracleParser;
import java.util.ArrayList;
@@ -43,6 +45,9 @@ public class OpenGaussGenerator {
else if (node instanceof JoinSourceTabNode) {
return GenJoinSQL(node);
}
+ else if (node instanceof UpdateNode) {
+ return GenUpdateSQL(node);
+ }
else {
try {
throw new GenerateFailedException("Root node:" + node.getClass() + "(Unsupported node type!)");
@@ -72,7 +77,7 @@ public class OpenGaussGenerator {
private String GenSelectSQL(ASTNode node) {
visitSelect(node, null);
- System.out.println(node.getASTString());
+// System.out.println(node.getASTString());
return node.toQueryString();
}
@@ -81,6 +86,12 @@ public class OpenGaussGenerator {
return node.toQueryString();
}
+ private String GenUpdateSQL(ASTNode node) {
+ visitUpdate(node, null);
+// System.out.println(node.getASTString());
+ return node.toQueryString();
+ }
+
private void visitCrt(ASTNode node) {
if (node instanceof ColumnNode) {
ColumnTypeConvert((ColumnNode) node);
@@ -138,6 +149,24 @@ public class OpenGaussGenerator {
}
}
+ private void visitUpdate(ASTNode node, ASTNode parentNode) {
+ if (node instanceof UpdateObjNode) {
+ // check whether there exists a join statement in the select sql
+ if (node.getTokens().contains(new Token(Token.TokenType.KEYWORD, "JOIN"))) {
+ ASTNode joinRootNode = OracleParser.parseJoin(node.getTokens());
+ visitJoin(joinRootNode);
+ parentNode.replaceChild(node, joinRootNode);
+ ASTNode smallestChild = joinRootNode.getDeepestChild();
+ for (ASTNode child : node.getChildren()) {
+ smallestChild.addChild(child);
+ }
+ }
+ }
+ for (ASTNode child : node.getChildren()) {
+ visitUpdate(child, node);
+ }
+ }
+
private void ColumnTypeConvert(ColumnNode node) {
// type convert
if (node.getType().getValue().equalsIgnoreCase("NUMBER")) {
diff --git a/sqlTranslate/src/main/java/Lexer/OracleLexer.java b/sqlTranslate/src/main/java/Lexer/OracleLexer.java
index 90b2554c3..39e2987c2 100644
--- a/sqlTranslate/src/main/java/Lexer/OracleLexer.java
+++ b/sqlTranslate/src/main/java/Lexer/OracleLexer.java
@@ -21,6 +21,8 @@ public class OracleLexer {
, "DISTINCT", "JOIN", "GROUP BY", "ORDER BY", "HAVING", "UNION", "CASE", "WHEN", "END", "AS"
// keywords of join
, "INNER JOIN", "LEFT JOIN", "LEFT OUTER JOIN", "RIGHT JOIN", "RIGHT OUTER JOIN", "FULL JOIN", "FULL OUTER JOIN", "CROSS JOIN", "JOIN", "ON", "USING"
+ // keywords of update
+ , "SET"
};
private static final Pattern TOKEN_PATTERN = Pattern.compile(
"(NUMBER\\(.*?\\))|" + // NUMBER() function
diff --git a/sqlTranslate/src/main/java/Main.java b/sqlTranslate/src/main/java/Main.java
index db3af17fd..b2381420a 100644
--- a/sqlTranslate/src/main/java/Main.java
+++ b/sqlTranslate/src/main/java/Main.java
@@ -18,6 +18,11 @@ public class Main {
// String sql = "DROP TABLE employees CASCADE CONSTRAINTS;";
String sql = "SELECT e.first_name, d.department_name FROM employees e JOIN departments d Using e.department_id = d.department_id;";
// String sql = "employees e JOIN departments d ON e.department_id = d.department_id;";
+// String sql = "UPDATE employees e\n" +
+// "JOIN departments d using e.department_id = d.department_id\n" +
+// "SET e.salary = e.salary * 1.10,\n" +
+// " d.budget = d.budget * 1.10\n" +
+// "WHERE d.department_name = 'Sales';";
OracleLexer lexer = new OracleLexer(sql);
lexer.printTokens();
OracleParser parser = new OracleParser(lexer);
diff --git a/sqlTranslate/src/main/java/Parser/AST/Join/JoinEndNode.java b/sqlTranslate/src/main/java/Parser/AST/Join/JoinEndNode.java
index 53ca8eda5..96e252c69 100644
--- a/sqlTranslate/src/main/java/Parser/AST/Join/JoinEndNode.java
+++ b/sqlTranslate/src/main/java/Parser/AST/Join/JoinEndNode.java
@@ -20,5 +20,9 @@ public class JoinEndNode extends ASTNode {
@Override
public void visit(ASTNode node, StringBuilder queryString) {
queryString.append(toString());
+ for (ASTNode child : getChildren())
+ {
+ child.visit(child, queryString);
+ }
}
}
diff --git a/sqlTranslate/src/main/java/Parser/AST/SelectStatementNode.java b/sqlTranslate/src/main/java/Parser/AST/SelectStatementNode.java
deleted file mode 100644
index 0d3eb8280..000000000
--- a/sqlTranslate/src/main/java/Parser/AST/SelectStatementNode.java
+++ /dev/null
@@ -1,27 +0,0 @@
-package Parser.AST;
-
-import Lexer.Token;
-
-import java.util.List;
-
-public class SelectStatementNode extends ASTNode{
- public SelectStatementNode(ASTNode node)
- {
- super(node);
- }
-
- public SelectStatementNode(List tokens)
- {
- super(tokens);
- }
-
- @Override
- public void visit(ASTNode node, StringBuilder queryString)
- {
- queryString.append(toString() + " ");
- for (ASTNode child : getChildren())
- {
- child.visit(child, queryString);
- }
- }
-}
diff --git a/sqlTranslate/src/main/java/Parser/OracleParser.java b/sqlTranslate/src/main/java/Parser/OracleParser.java
index 31d5be422..312bb14d2 100644
--- a/sqlTranslate/src/main/java/Parser/OracleParser.java
+++ b/sqlTranslate/src/main/java/Parser/OracleParser.java
@@ -19,6 +19,7 @@ import Parser.AST.Insert.InsertNode;
import Parser.AST.Insert.InsertObjNode;
import Parser.AST.Join.*;
import Parser.AST.Select.*;
+import Parser.AST.Update.*;
import java.util.Stack;
@@ -46,6 +47,9 @@ public class OracleParser {
else if (lexer.getTokens().get(0).getValue().equalsIgnoreCase("SELECT")) {
return parseSelect(lexer.getTokens());
}
+ else if (lexer.getTokens().get(0).getValue().equalsIgnoreCase("UPDATE")) {
+ return parseUpdate(lexer.getTokens());
+ }
else {
try {
throw new ParseFailedException("Parse failed!");
@@ -739,4 +743,79 @@ public class OracleParser {
return root;
}
+
+ /**
+ * Update
+ * For example: UPDATE table_name SET column1 = value1, column2 = value2, ... WHERE condition;
+ */
+ private ASTNode parseUpdate(List parseTokens) {
+ List tokens = new ArrayList<>();
+ tokens.add(parseTokens.get(0));
+ ASTNode root = new UpdateNode(tokens);
+ ASTNode currentNode = root;
+ for (int i = 1; i < parseTokens.size(); i++) {
+ if (i == 1 && parseTokens.get(i).hasType(Token.TokenType.IDENTIFIER)) {
+ // match update object
+ tokens = new ArrayList<>();
+ for (int j = i; j < parseTokens.size(); j++) {
+ if (parseTokens.get(j).hasType(Token.TokenType.KEYWORD) && parseTokens.get(j).getValue().equalsIgnoreCase("SET")) {
+ i = j - 1;
+ break;
+ }
+ tokens.add(parseTokens.get(j));
+ }
+ ASTNode childNode = new UpdateObjNode(tokens);
+ currentNode.addChild(childNode);
+ currentNode = childNode;
+ }
+ else if (parseTokens.get(i).hasType(Token.TokenType.KEYWORD) && parseTokens.get(i).getValue().equalsIgnoreCase("SET")) {
+ // match SET
+ tokens = new ArrayList<>();
+ for (int j = i; j < parseTokens.size(); j++) {
+ if (parseTokens.get(j).hasType(Token.TokenType.KEYWORD) && parseTokens.get(j).getValue().equalsIgnoreCase("WHERE")) {
+ i = j - 1;
+ break;
+ }
+ tokens.add(parseTokens.get(j));
+ }
+ ASTNode childNode = new UpdateSetNode(tokens);
+ currentNode.addChild(childNode);
+ currentNode = childNode;
+ }
+ else if (parseTokens.get(i).hasType(Token.TokenType.KEYWORD) && parseTokens.get(i).getValue().equalsIgnoreCase("WHERE")) {
+ // match WHERE
+ tokens = new ArrayList<>();
+ for (int j = i; j < parseTokens.size(); j++) {
+ if (parseTokens.get(j).hasType(Token.TokenType.SYMBOL) && parseTokens.get(j).getValue().equalsIgnoreCase(";")) {
+ i = j - 1;
+ break;
+ }
+ tokens.add(parseTokens.get(j));
+ }
+ ASTNode childNode = new UpdateWhereNode(tokens);
+ currentNode.addChild(childNode);
+ currentNode = childNode;
+ }
+ else if (parseTokens.get(i).hasType(Token.TokenType.SYMBOL) && parseTokens.get(i).getValue().equalsIgnoreCase(";")) {
+ tokens = new ArrayList<>();
+ tokens.add(parseTokens.get(i));
+ ASTNode childNode = new UpdateEndNode(tokens);
+ currentNode.addChild(childNode);
+ currentNode = childNode;
+ }
+ else if (parseTokens.get(i).hasType(Token.TokenType.EOF)) {
+ break;
+ }
+ else {
+ try {
+ throw new ParseFailedException("Parse failed:" + parseTokens.get(i));
+ }
+ catch (ParseFailedException e) {
+ e.printStackTrace();
+ }
+ }
+ }
+
+ return root;
+ }
}