#1261 - ENH: Support use of getGeneratedKeys with SqlUpdate

This commit is contained in:
Rob Bygrave
2018-02-16 09:45:49 +13:00
parent 9926cad7f1
commit 5237e8ba66
12 changed files with 315 additions and 13 deletions
+24
View File
@@ -54,6 +54,25 @@ public interface SqlUpdate {
*/
int execute();
/**
* Return the generated key value.
*/
Object getGeneratedKey();
/**
* Execute and return the generated key. This is effectively a short cut for:
*
* <pre>{@code
*
* sqlUpdate.execute();
* Object key = sqlUpdate.getGeneratedKey();
*
* }</pre>
*
* @return The generated key value
*/
Object executeGetKey();
/**
* Return true if eBean should automatically deduce the table modification
* information and process it.
@@ -88,6 +107,11 @@ public interface SqlUpdate {
*/
SqlUpdate setLabel(String label);
/**
* Set to true when we want to use getGeneratedKeys with this statement.
*/
SqlUpdate setGetGeneratedKeys(boolean getGeneratedKeys);
/**
* Return the sql statement.
*/
@@ -4,7 +4,23 @@ import io.ebean.SqlUpdate;
public interface SpiSqlUpdate extends SqlUpdate {
/**
* Return the Bind parameters.
*/
BindParams getBindParams();
/**
* Set the final sql being executed with named parameters replaced etc.
*/
void setGeneratedSql(String sql);
/**
* Return true if we are using getGeneratedKeys.
*/
boolean isGetGeneratedKeys();
/**
* Set the generated key value.
*/
void setGeneratedKey(Object idValue);
}
@@ -65,6 +65,10 @@ public final class DefaultSqlUpdate implements Serializable, SpiSqlUpdate {
*/
private int addPos;
private boolean getGeneratedKeys;
private Object generatedKey;
/**
* Create with server sql and bindParams object.
* <p>
@@ -93,6 +97,12 @@ public final class DefaultSqlUpdate implements Serializable, SpiSqlUpdate {
this(null, sql, new BindParams());
}
@Override
public Object executeGetKey() {
execute();
return getGeneratedKey();
}
@Override
public int execute() {
if (server != null) {
@@ -103,6 +113,16 @@ public final class DefaultSqlUpdate implements Serializable, SpiSqlUpdate {
}
}
@Override
public Object getGeneratedKey() {
return generatedKey;
}
@Override
public void setGeneratedKey(Object idValue) {
this.generatedKey = idValue;
}
@Override
public boolean isAutoTableMod() {
return isAutoTableMod;
@@ -125,6 +145,17 @@ public final class DefaultSqlUpdate implements Serializable, SpiSqlUpdate {
return this;
}
@Override
public boolean isGetGeneratedKeys() {
return getGeneratedKeys;
}
@Override
public SqlUpdate setGetGeneratedKeys(boolean getGeneratedKeys) {
this.getGeneratedKeys = getGeneratedKeys;
return this;
}
@Override
public String getGeneratedSql() {
return generatedSql;
@@ -73,6 +73,11 @@ public final class PersistRequestUpdateSql extends PersistRequest {
*/
@Override
public void setGeneratedKey(Object idValue) {
updateSql.setGeneratedKey(idValue);
}
public boolean isGetGeneratedKeys() {
return updateSql.isGetGeneratedKeys();
}
/**
@@ -16,7 +16,7 @@ import java.sql.SQLException;
/**
* Executes the UpdateSql requests.
*/
public class ExeOrmUpdate {
class ExeOrmUpdate {
private static final Logger logger = LoggerFactory.getLogger(ExeOrmUpdate.class);
@@ -27,7 +27,7 @@ public class ExeOrmUpdate {
/**
* Create with a given binder.
*/
public ExeOrmUpdate(Binder binder) {
ExeOrmUpdate(Binder binder) {
this.pstmtFactory = new PstmtFactory();
this.binder = binder;
}
@@ -107,7 +107,7 @@ public class ExeOrmUpdate {
if (logSql) {
t.logSql(sql);
}
pstmt = pstmtFactory.getPstmt(t, sql);
pstmt = pstmtFactory.getPstmt(t, sql, false);
}
String bindLog = null;
@@ -6,17 +6,19 @@ import io.ebeaninternal.api.SpiTransaction;
import io.ebeaninternal.server.core.PersistRequestUpdateSql;
import io.ebeaninternal.server.core.PersistRequestUpdateSql.SqlType;
import io.ebeaninternal.server.util.BindParamsParser;
import io.ebeaninternal.util.JdbcClose;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import javax.persistence.PersistenceException;
import java.sql.PreparedStatement;
import java.sql.ResultSet;
import java.sql.SQLException;
/**
* Executes the UpdateSql requests.
*/
public class ExeUpdateSql {
class ExeUpdateSql {
private static final Logger logger = LoggerFactory.getLogger(ExeUpdateSql.class);
@@ -27,7 +29,7 @@ public class ExeUpdateSql {
/**
* Create with a given binder.
*/
public ExeUpdateSql(Binder binder) {
ExeUpdateSql(Binder binder) {
this.binder = binder;
this.pstmtFactory = new PstmtFactory();
}
@@ -51,6 +53,9 @@ public class ExeUpdateSql {
} else {
int rowCount = pstmt.executeUpdate();
request.checkRowCount(rowCount);
if (request.isGetGeneratedKeys()) {
readGeneratedKeys(pstmt, request);
}
request.postExecute();
return rowCount;
}
@@ -68,6 +73,22 @@ public class ExeUpdateSql {
}
}
private void readGeneratedKeys(PreparedStatement stmt, PersistRequestUpdateSql request) {
ResultSet resultSet = null;
try {
resultSet = stmt.getGeneratedKeys();
if (resultSet.next()) {
request.setGeneratedKey(resultSet.getObject(1));
}
} catch (SQLException ex) {
throw new PersistenceException(ex);
} finally {
JdbcClose.close(resultSet);
}
}
private PreparedStatement bindStmt(PersistRequestUpdateSql request, boolean batchThisRequest) throws SQLException {
SpiSqlUpdate updateSql = request.getUpdateSql();
@@ -90,7 +111,7 @@ public class ExeUpdateSql {
if (logSql) {
t.logSql(TrimLogSql.trim(sql));
}
pstmt = pstmtFactory.getPstmt(t, sql);
pstmt = pstmtFactory.getPstmt(t, sql, request.isGetGeneratedKeys());
}
if (updateSql.getTimeout() > 0) {
@@ -6,6 +6,7 @@ import java.sql.CallableStatement;
import java.sql.Connection;
import java.sql.PreparedStatement;
import java.sql.SQLException;
import java.sql.Statement;
/**
* Factory for creating Statements.
@@ -14,9 +15,9 @@ import java.sql.SQLException;
* getGeneratedKeys.
* </p>
*/
public class PstmtFactory {
class PstmtFactory {
public PstmtFactory() {
PstmtFactory() {
}
/**
@@ -30,9 +31,13 @@ public class PstmtFactory {
/**
* Get a prepared statement without any batching.
*/
public PreparedStatement getPstmt(SpiTransaction t, String sql) throws SQLException {
public PreparedStatement getPstmt(SpiTransaction t, String sql, boolean getGeneratedKeys) throws SQLException {
Connection conn = t.getInternalConnection();
return conn.prepareStatement(sql);
if (getGeneratedKeys) {
return conn.prepareStatement(sql, Statement.RETURN_GENERATED_KEYS);
} else {
return conn.prepareStatement(sql);
}
}
/**
@@ -24,6 +24,8 @@ public abstract class DmlHandler implements PersistHandler, BindableRequest {
private static final Logger logger = LoggerFactory.getLogger(DmlHandler.class);
private static final int[] GENERATED_KEY_COLUMNS = new int[]{1};
/**
* The originating request.
*/
@@ -260,8 +262,7 @@ public abstract class DmlHandler implements PersistHandler, BindableRequest {
// the Id generated is always the first column
// Required to stop Oracle10 giving us Oracle rowId??
// Other jdbc drivers seem fine without this hint.
int[] columns = {1};
return conn.prepareStatement(sql, columns);
return conn.prepareStatement(sql, GENERATED_KEY_COLUMNS);
} else {
return conn.prepareStatement(sql);
@@ -4,6 +4,7 @@ import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import java.sql.Connection;
import java.sql.ResultSet;
import java.sql.SQLException;
/**
@@ -13,6 +14,19 @@ public class JdbcClose {
private static final Logger logger = LoggerFactory.getLogger(JdbcClose.class);
/**
* Close the resultSet logging if an error occurs.
*/
public static void close(ResultSet resultSet) {
try {
if (resultSet != null) {
resultSet.close();
}
} catch (SQLException e) {
logger.warn("Error closing resultSet", e);
}
}
/**
* Close the connection logging if an error occurs.
*/