@@ -112,6 +112,8 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
|
||||
private @Nullable ReactiveEntityCallbacks entityCallbacks;
|
||||
|
||||
private Function<Statement, Statement> statementFilterFunction = Function.identity();
|
||||
|
||||
/**
|
||||
* Create a new {@link R2dbcEntityTemplate} given {@link ConnectionFactory}.
|
||||
*
|
||||
@@ -174,6 +176,19 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
this.projectionFactory = new SpelAwareProxyProjectionFactory();
|
||||
}
|
||||
|
||||
/**
|
||||
* Set a {@link Function Statement Filter Function} that is applied to every {@link Statement}.
|
||||
*
|
||||
* @param statementFilterFunction must not be {@literal null}.
|
||||
* @since 3.4
|
||||
*/
|
||||
public void setStatementFilterFunction(Function<Statement, Statement> statementFilterFunction) {
|
||||
|
||||
Assert.notNull(statementFilterFunction, "StatementFilterFunction must not be null");
|
||||
|
||||
this.statementFilterFunction = statementFilterFunction;
|
||||
}
|
||||
|
||||
@Override
|
||||
public DatabaseClient getDatabaseClient() {
|
||||
return this.databaseClient;
|
||||
@@ -274,6 +289,7 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
PreparedOperation<?> operation = statementMapper.getMappedObject(selectSpec);
|
||||
|
||||
return this.databaseClient.sql(operation) //
|
||||
.filter(statementFilterFunction) //
|
||||
.map((r, md) -> r.get(0, Long.class)) //
|
||||
.first() //
|
||||
.defaultIfEmpty(0L);
|
||||
@@ -302,6 +318,7 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
PreparedOperation<?> operation = statementMapper.getMappedObject(selectSpec);
|
||||
|
||||
return this.databaseClient.sql(operation) //
|
||||
.filter(statementFilterFunction) //
|
||||
.map((r, md) -> r) //
|
||||
.first() //
|
||||
.hasElement();
|
||||
@@ -362,7 +379,7 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
PreparedOperation<?> operation = statementMapper.getMappedObject(selectSpec);
|
||||
|
||||
return getRowsFetchSpec(
|
||||
databaseClient.sql(operation).filter(filterFunction),
|
||||
databaseClient.sql(operation).filter(statementFilterFunction.andThen(filterFunction)),
|
||||
entityType,
|
||||
returnType
|
||||
);
|
||||
@@ -397,7 +414,7 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
}
|
||||
|
||||
PreparedOperation<?> operation = statementMapper.getMappedObject(selectSpec);
|
||||
return this.databaseClient.sql(operation).fetch().rowsUpdated();
|
||||
return this.databaseClient.sql(operation).filter(statementFilterFunction).fetch().rowsUpdated();
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -422,7 +439,7 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
}
|
||||
|
||||
PreparedOperation<?> operation = statementMapper.getMappedObject(deleteSpec);
|
||||
return this.databaseClient.sql(operation).fetch().rowsUpdated().defaultIfEmpty(0L);
|
||||
return this.databaseClient.sql(operation).filter(statementFilterFunction).fetch().rowsUpdated().defaultIfEmpty(0L);
|
||||
}
|
||||
|
||||
// -------------------------------------------------------------------------
|
||||
@@ -441,7 +458,8 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
Assert.notNull(operation, "PreparedOperation must not be null");
|
||||
Assert.notNull(entityClass, "Entity class must not be null");
|
||||
|
||||
return new EntityCallbackAdapter<>(getRowsFetchSpec(databaseClient.sql(operation), entityClass, resultType),
|
||||
return new EntityCallbackAdapter<>(
|
||||
getRowsFetchSpec(databaseClient.sql(operation).filter(statementFilterFunction), entityClass, resultType),
|
||||
getTableNameOrEmpty(entityClass));
|
||||
}
|
||||
|
||||
@@ -451,7 +469,8 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
Assert.notNull(operation, "PreparedOperation must not be null");
|
||||
Assert.notNull(rowMapper, "Row mapper must not be null");
|
||||
|
||||
return new EntityCallbackAdapter<>(databaseClient.sql(operation).map(rowMapper), SqlIdentifier.EMPTY);
|
||||
return new EntityCallbackAdapter<>(databaseClient.sql(operation).filter(statementFilterFunction).map(rowMapper),
|
||||
SqlIdentifier.EMPTY);
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -462,7 +481,8 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
Assert.notNull(entityClass, "Entity class must not be null");
|
||||
Assert.notNull(rowMapper, "Row mapper must not be null");
|
||||
|
||||
return new EntityCallbackAdapter<>(databaseClient.sql(operation).map(rowMapper), getTableNameOrEmpty(entityClass));
|
||||
return new EntityCallbackAdapter<>(databaseClient.sql(operation).filter(statementFilterFunction).map(rowMapper),
|
||||
getTableNameOrEmpty(entityClass));
|
||||
}
|
||||
|
||||
// -------------------------------------------------------------------------
|
||||
@@ -541,6 +561,8 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
return this.databaseClient.sql(operation) //
|
||||
.filter(statement -> {
|
||||
|
||||
statement = statementFilterFunction.apply(statement);
|
||||
|
||||
if (identifierColumns.isEmpty()) {
|
||||
return statement.returnGeneratedValues();
|
||||
}
|
||||
@@ -632,6 +654,7 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw
|
||||
PreparedOperation<?> operation = mapper.getMappedObject(updateSpec);
|
||||
|
||||
return this.databaseClient.sql(operation) //
|
||||
.filter(statementFilterFunction) //
|
||||
.fetch() //
|
||||
.rowsUpdated() //
|
||||
.handle((rowsUpdated, sink) -> {
|
||||
|
||||
@@ -206,8 +206,6 @@ public class R2dbcEntityTemplateUnitTests {
|
||||
@Test // GH-469
|
||||
void shouldExistsByCriteria() {
|
||||
|
||||
MockRowMetadata metadata = MockRowMetadata.builder()
|
||||
.columnMetadata(MockColumnMetadata.builder().name("name").type(R2dbcType.VARCHAR).build()).build();
|
||||
MockResult result = MockResult.builder().row(MockRow.builder().identified(0, Long.class, 1L).build()).build();
|
||||
|
||||
recorder.addStubbing(s -> s.startsWith("SELECT"), result);
|
||||
@@ -654,6 +652,24 @@ public class R2dbcEntityTemplateUnitTests {
|
||||
}).verifyComplete();
|
||||
}
|
||||
|
||||
@Test // GH-1652
|
||||
void shouldConsiderFilterFunction() {
|
||||
|
||||
MockResult result = MockResult.builder().row(MockRow.builder().identified(0, Long.class, 1L).build()).build();
|
||||
|
||||
recorder.addStubbing(s -> s.startsWith("SELECT"), result);
|
||||
|
||||
entityTemplate.setStatementFilterFunction(statement -> statement.fetchSize(10));
|
||||
entityTemplate.count(Query.empty(), Person.class) //
|
||||
.as(StepVerifier::create) //
|
||||
.expectNext(1L) //
|
||||
.verifyComplete();
|
||||
|
||||
StatementRecorder.RecordedStatement statement = recorder.getCreatedStatement(s -> s.startsWith("SELECT"));
|
||||
|
||||
assertThat(statement.getFetchSize()).isEqualTo(10);
|
||||
}
|
||||
|
||||
@ReadingConverter
|
||||
static class PkConverter implements Converter<ByteBuffer, DoubleHolder> {
|
||||
|
||||
|
||||
Reference in New Issue
Block a user