diff --git a/docs/src/main/asciidoc/spanner.adoc b/docs/src/main/asciidoc/spanner.adoc index 492cb7a06a4..cedec0501b2 100644 --- a/docs/src/main/asciidoc/spanner.adoc +++ b/docs/src/main/asciidoc/spanner.adoc @@ -1081,6 +1081,42 @@ Parameters can also be of type `Struct` or POJOs. If a POJO is given as a parameter, it will be converted to a `Struct` with the same type-conversion logic as used to create write mutations. Comparisons using Struct parameters are limited to https://cloud.google.com/spanner/docs/data-types#limited-comparisons-for-struct[what is available with Cloud Spanner]. +==== Locking rows with `FOR UPDATE` + +For workloads with high write contention, a query can acquire exclusive locks by setting +`forUpdate = true` on the `@Query` annotation. +The annotation can be used with a derived query without specifying a SQL string: + +[source,java] +---- +public interface TradeRepository extends SpannerRepository { + + @Query(forUpdate = true) + Optional findBySymbol(String symbol); +} +---- + +To retrieve an entity by its ID with exclusive locks, use `findByIdForUpdate`: + +[source,java] +---- +@Transactional(transactionManager = "spannerTransactionManager") +public void updateTrade(Key tradeId) { + Trade trade = tradeRepository.findByIdForUpdate(tradeId).orElseThrow(); + trade.setAction("SELL"); + tradeRepository.save(trade); +} +---- + +Both forms must be executed in a read-write transaction. +The generated queries also apply `FOR UPDATE` when loading eager or lazy interleaved children. +For programmatic queries with `SpannerTemplate`, set `forUpdate` on `SpannerQueryOptions`. + +`FOR UPDATE` can reduce transaction aborts when concurrent transactions read and update the same +data, but conflicting transactions wait for the locks and can therefore reduce throughput. +See https://cloud.google.com/spanner/docs/use-select-for-update[Use SELECT FOR UPDATE] for details +and restrictions. + ==== Query methods by convention diff --git a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/core/SpannerQueryOptions.java b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/core/SpannerQueryOptions.java index 6c930a613c3..b8052ff9c02 100644 --- a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/core/SpannerQueryOptions.java +++ b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/core/SpannerQueryOptions.java @@ -30,6 +30,8 @@ */ public class SpannerQueryOptions extends AbstractSpannerRequestOptions { + private boolean forUpdate; + /** * Constructor to create an instance. Use the extension-style add/set functions to add options and * settings. @@ -44,6 +46,22 @@ public SpannerQueryOptions addQueryOption(QueryOption queryOption) { return this; } + public boolean isForUpdate() { + return this.forUpdate; + } + + /** + * Sets whether the query should acquire exclusive locks on the selected rows. Cloud Spanner only + * supports {@code FOR UPDATE} in read-write transactions. + * + * @param forUpdate whether {@code FOR UPDATE} is enabled + * @return this options instance + */ + public SpannerQueryOptions setForUpdate(boolean forUpdate) { + this.forUpdate = forUpdate; + return this; + } + @Override public SpannerQueryOptions setIncludeProperties(Set includeProperties) { super.setIncludeProperties(includeProperties); diff --git a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/core/SpannerTemplate.java b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/core/SpannerTemplate.java index f8a63fbcf34..bdffb31d096 100644 --- a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/core/SpannerTemplate.java +++ b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/core/SpannerTemplate.java @@ -291,8 +291,10 @@ public List queryAll(Class entityClass, SpannerPageableQueryOptions op return query( entityClass, SpannerStatementQueryExecutor.buildStatementFromSqlWithArgs( - SpannerStatementQueryExecutor.applySortingPagingQueryOptions( - entityClass, options, sql, this.mappingContext, false), + SpannerStatementQueryExecutor.applyForUpdate( + SpannerStatementQueryExecutor.applySortingPagingQueryOptions( + entityClass, options, sql, this.mappingContext, false), + options.isForUpdate()), null, null, null, @@ -502,6 +504,10 @@ public T performReadOnlyTransaction( } public ResultSet executeQuery(Statement statement, SpannerQueryOptions options) { + if (options != null && options.isForUpdate() && !isReadWriteTransactionActive()) { + throw new SpannerDataException( + "FOR UPDATE queries must be executed in a read-write transaction."); + } ResultSet resultSet = performQuery(statement, options); if (LOGGER.isDebugEnabled()) { String message; @@ -640,7 +646,8 @@ private List queryAndResolveChildren( executeQuery(statement, options), entityClass, (options != null) ? options.getIncludeProperties() : null, - options != null && options.isAllowPartialRead()); + options != null && options.isAllowPartialRead(), + options != null && options.isForUpdate()); } private List mapToListAndResolveChildren( @@ -648,20 +655,36 @@ private List mapToListAndResolveChildren( Class entityClass, Set includeProperties, boolean allowMissingColumns) { + return mapToListAndResolveChildren( + resultSet, entityClass, includeProperties, allowMissingColumns, false); + } + + private List mapToListAndResolveChildren( + ResultSet resultSet, + Class entityClass, + Set includeProperties, + boolean allowMissingColumns, + boolean forUpdate) { return resolveChildEntities( this.spannerEntityProcessor.mapToList( resultSet, entityClass, includeProperties, allowMissingColumns), - includeProperties); + includeProperties, + forUpdate); } private List resolveChildEntities(List entities, Set includeProperties) { + return resolveChildEntities(entities, includeProperties, false); + } + + private List resolveChildEntities( + List entities, Set includeProperties, boolean forUpdate) { for (Object entity : entities) { - resolveChildEntity(entity, includeProperties); + resolveChildEntity(entity, includeProperties, forUpdate); } return entities; } - private void resolveChildEntity(Object entity, Set includeProperties) { + private void resolveChildEntity(Object entity, Set includeProperties, boolean forUpdate) { SpannerPersistentEntity spannerPersistentEntity = this.mappingContext.getPersistentEntityOrFail(entity.getClass()); @@ -675,7 +698,7 @@ private void resolveChildEntity(Object entity, Set includeProperties) { // an interleaved property can only be List List propertyValue = (List) accessor.getProperty(spannerPersistentProperty); if (propertyValue != null) { - resolveChildEntities(propertyValue, null); + resolveChildEntities(propertyValue, null, forUpdate); return; } Class childType = spannerPersistentProperty.getColumnInnerType(); @@ -688,8 +711,9 @@ private void resolveChildEntity(Object entity, Set includeProperties) { this.spannerSchemaUtils.getKey(entity), spannerPersistentProperty, this.spannerEntityProcessor.getWriteConverter(), - this.mappingContext), - null); + this.mappingContext, + forUpdate), + new SpannerQueryOptions().setForUpdate(forUpdate)); accessor.setProperty( spannerPersistentProperty, @@ -718,6 +742,10 @@ private TransactionContext getTransactionContext() { return null; } + private boolean isReadWriteTransactionActive() { + return this instanceof ReadWriteTransactionSpannerTemplate || getTransactionContext() != null; + } + private A doWithOrWithoutTransactionContext( Function funcWithTransactionContext, Supplier funcWithoutTransactionContext) { diff --git a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/SpannerRepository.java b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/SpannerRepository.java index 56b441637a5..938be5e3f2e 100644 --- a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/SpannerRepository.java +++ b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/SpannerRepository.java @@ -17,6 +17,7 @@ package com.google.cloud.spring.data.spanner.repository; import com.google.cloud.spring.data.spanner.core.SpannerOperations; +import java.util.Optional; import java.util.function.Function; import org.springframework.data.repository.CrudRepository; import org.springframework.data.repository.PagingAndSortingRepository; @@ -58,4 +59,13 @@ public interface SpannerRepository * @return the final result of the transaction. */ A performReadOnlyTransaction(Function, A> operations); + + /** + * Retrieves an entity by its id and acquires exclusive locks on the selected row and its + * interleaved children. This method must be called in a read-write transaction. + * + * @param id the entity id + * @return the entity, or empty if none was found + */ + Optional findByIdForUpdate(I id); } diff --git a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/PartTreeSpannerQuery.java b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/PartTreeSpannerQuery.java index 279f0d4ebcb..18280a94853 100644 --- a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/PartTreeSpannerQuery.java +++ b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/PartTreeSpannerQuery.java @@ -65,7 +65,8 @@ protected List executeRawResult(Object[] parameters) { paramAccessor, getQueryMethod().getQueryMethod().getParameters(), this.spannerTemplate, - this.spannerMappingContext); + this.spannerMappingContext, + false); } if (this.tree.isDelete()) { return this.spannerTemplate.performReadWriteTransaction(getDeleteFunction(parameters)); @@ -76,7 +77,8 @@ protected List executeRawResult(Object[] parameters) { paramAccessor, getQueryMethod().getQueryMethod().getParameters(), this.spannerTemplate, - this.spannerMappingContext); + this.spannerMappingContext, + this.queryMethod.isForUpdate()); } private Function getDeleteFunction(Object[] parameters) { @@ -90,7 +92,8 @@ private Function getDeleteFunction(Object[] parameters) { paramAccessor, getQueryMethod().getQueryMethod().getParameters(), transactionTemplate, - this.spannerMappingContext); + this.spannerMappingContext, + false); transactionTemplate.deleteAll(entitiesToDelete); List result = null; diff --git a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/Query.java b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/Query.java index 4b989203843..28045ad3646 100644 --- a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/Query.java +++ b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/Query.java @@ -53,4 +53,12 @@ * method is executed as a DML query. */ boolean dmlStatement() default false; + + /** + * Indicates if the query should acquire exclusive locks on the selected rows. This option is + * valid only for select queries executed in a read-write transaction. + * + * @return {@code true} if {@code FOR UPDATE} should be applied. + */ + boolean forUpdate() default false; } diff --git a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SpannerQueryMethod.java b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SpannerQueryMethod.java index 754835d87a7..590e9b6ba08 100644 --- a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SpannerQueryMethod.java +++ b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SpannerQueryMethod.java @@ -85,4 +85,9 @@ Method getQueryMethod() { Query getQueryAnnotation() { return AnnotatedElementUtils.findMergedAnnotation(this.queryMethod, Query.class); } + + boolean isForUpdate() { + Query query = getQueryAnnotation(); + return query != null && query.forUpdate(); + } } diff --git a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SpannerStatementQueryExecutor.java b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SpannerStatementQueryExecutor.java index b4d38bd3837..cd0962c275b 100644 --- a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SpannerStatementQueryExecutor.java +++ b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SpannerStatementQueryExecutor.java @@ -22,6 +22,7 @@ import com.google.cloud.spanner.Struct; import com.google.cloud.spanner.ValueBinder; import com.google.cloud.spring.data.spanner.core.SpannerPageableQueryOptions; +import com.google.cloud.spring.data.spanner.core.SpannerQueryOptions; import com.google.cloud.spring.data.spanner.core.SpannerTemplate; import com.google.cloud.spring.data.spanner.core.convert.ConversionUtils; import com.google.cloud.spring.data.spanner.core.convert.ConverterAwareMappingSpannerEntityWriter; @@ -87,8 +88,26 @@ public static List executeQuery( Parameter[] queryMethodParamsMetadata, SpannerTemplate spannerTemplate, SpannerMappingContext spannerMappingContext) { + return executeQuery( + type, + tree, + parameterAccessor, + queryMethodParamsMetadata, + spannerTemplate, + spannerMappingContext, + false); + } + + public static List executeQuery( + Class type, + PartTree tree, + ParameterAccessor parameterAccessor, + Parameter[] queryMethodParamsMetadata, + SpannerTemplate spannerTemplate, + SpannerMappingContext spannerMappingContext, + boolean forUpdate) { SqlStringAndPlaceholders sqlStringAndPlaceholders = - buildPartTreeSqlString(tree, spannerMappingContext, type, parameterAccessor); + buildPartTreeSqlString(tree, spannerMappingContext, type, parameterAccessor, forUpdate); Map paramMetadataMap = preparePartTreeSqlTagParameterMap(queryMethodParamsMetadata, sqlStringAndPlaceholders); Object[] params = StreamSupport.stream(parameterAccessor.spliterator(), false).toArray(); @@ -101,7 +120,7 @@ public static List executeQuery( spannerTemplate.getSpannerEntityProcessor().getWriteConverter(), params, paramMetadataMap), - null); + new SpannerQueryOptions().setForUpdate(forUpdate)); } private static Map preparePartTreeSqlTagParameterMap( @@ -143,8 +162,28 @@ public static List executeQuery( Parameter[] queryMethodParamsMetadata, SpannerTemplate spannerTemplate, SpannerMappingContext spannerMappingContext) { + return executeQuery( + rowFunc, + type, + tree, + parameterAccessor, + queryMethodParamsMetadata, + spannerTemplate, + spannerMappingContext, + false); + } + + public static List executeQuery( + Function rowFunc, + Class type, + PartTree tree, + ParameterAccessor parameterAccessor, + Parameter[] queryMethodParamsMetadata, + SpannerTemplate spannerTemplate, + SpannerMappingContext spannerMappingContext, + boolean forUpdate) { SqlStringAndPlaceholders sqlStringAndPlaceholders = - buildPartTreeSqlString(tree, spannerMappingContext, type, parameterAccessor); + buildPartTreeSqlString(tree, spannerMappingContext, type, parameterAccessor, forUpdate); Map paramMetadataMap = preparePartTreeSqlTagParameterMap(queryMethodParamsMetadata, sqlStringAndPlaceholders); Object[] params = StreamSupport.stream(parameterAccessor.spliterator(), false).toArray(); @@ -157,7 +196,22 @@ public static List executeQuery( spannerTemplate.getSpannerEntityProcessor().getWriteConverter(), params, paramMetadataMap), - null); + new SpannerQueryOptions().setForUpdate(forUpdate)); + } + + /** + * Appends {@code FOR UPDATE} to a select query when requested. + * + * @param sql SQL query + * @param forUpdate whether locking is requested + * @return the resulting SQL query + */ + public static String applyForUpdate(String sql, boolean forUpdate) { + return forUpdate + && !sql.regionMatches( + true, sql.length() - "FOR UPDATE".length(), "FOR UPDATE", 0, "FOR UPDATE".length()) + ? sql + " FOR UPDATE" + : sql; } /** @@ -196,7 +250,9 @@ public static String applySortingPagingQueryOptions( mappingContext.getPersistentEntityOrFail(entityClass); final String subquery = - fetchInterleaved ? getChildrenSubquery(persistentEntity, mappingContext) : ""; + fetchInterleaved + ? getChildrenSubquery(persistentEntity, mappingContext, options.isForUpdate()) + : ""; final String alias = subquery.isEmpty() ? "" : " " + persistentEntity.tableName(); StringBuilder sb = applySort( @@ -248,12 +304,28 @@ public static Statement getChildrenRowsQuery( SpannerPersistentProperty spannerPersistentProperty, SpannerCustomConverter writeConverter, SpannerMappingContext mappingContext) { + return getChildrenRowsQuery( + parentKey, spannerPersistentProperty, writeConverter, mappingContext, false); + } + + public static Statement getChildrenRowsQuery( + Key parentKey, + SpannerPersistentProperty spannerPersistentProperty, + SpannerCustomConverter writeConverter, + SpannerMappingContext mappingContext, + boolean forUpdate) { Class childType = spannerPersistentProperty.getColumnInnerType(); SpannerPersistentEntity persistentEntity = mappingContext.getPersistentEntityOrFail(childType); String whereClause = getWhere(spannerPersistentProperty, persistentEntity); return buildQuery( - KeySet.singleKey(parentKey), persistentEntity, writeConverter, mappingContext, whereClause); + KeySet.singleKey(parentKey), + persistentEntity, + writeConverter, + mappingContext, + whereClause, + null, + forUpdate); } /** @@ -298,7 +370,8 @@ public static Statement buildQuery( SpannerCustomConverter writeConverter, SpannerMappingContext mappingContext, String whereClause) { - return buildQuery(keySet, persistentEntity, writeConverter, mappingContext, whereClause, null); + return buildQuery( + keySet, persistentEntity, writeConverter, mappingContext, whereClause, null, false); } /** @@ -325,6 +398,18 @@ public static Statement buildQuery( SpannerMappingContext mappingContext, String whereClause, String index) { + return buildQuery( + keySet, persistentEntity, writeConverter, mappingContext, whereClause, index, false); + } + + public static Statement buildQuery( + KeySet keySet, + SpannerPersistentEntity persistentEntity, + SpannerCustomConverter writeConverter, + SpannerMappingContext mappingContext, + String whereClause, + String index, + boolean forUpdate) { List orParts = new ArrayList<>(); List tags = new ArrayList<>(); List keyParts = new ArrayList(); @@ -349,12 +434,13 @@ public static Statement buildQuery( String condition = combineWithAnd(keyClause, whereClause); String sb = "SELECT " - + getColumnsStringForSelect(persistentEntity, mappingContext, true) + + getColumnsStringForSelect(persistentEntity, mappingContext, true, forUpdate) + " FROM " + (!StringUtils.hasLength(index) ? persistentEntity.tableName() : String.format("%s@{FORCE_INDEX=%s}", persistentEntity.tableName(), index)) - + (condition.isEmpty() ? "" : WHERE + condition); + + (condition.isEmpty() ? "" : WHERE + condition) + + (forUpdate ? " FOR UPDATE" : ""); return buildStatementFromSqlWithArgs(sb, tags, null, writeConverter, keyParts.toArray(), null); } @@ -363,7 +449,8 @@ private static String getChildrenStructsQuery( SpannerPersistentEntity

parentPersistentEntity, SpannerMappingContext mappingContext, String columnName, - String whereClause) { + String whereClause, + boolean forUpdate) { String tableName = childPersistentEntity.tableName(); List parentKeyProperties = parentPersistentEntity.getFlattenedPrimaryKeyProperties(); @@ -386,6 +473,7 @@ private static String getChildrenStructsQuery( + tableName + WHERE + condition + + (forUpdate ? " FOR UPDATE" : "") + ") AS " + columnName; } @@ -495,9 +583,18 @@ public static String getColumnsStringForSelect( SpannerPersistentEntity spannerPersistentEntity, SpannerMappingContext mappingContext, boolean fetchInterleaved) { + return getColumnsStringForSelect( + spannerPersistentEntity, mappingContext, fetchInterleaved, false); + } + + public static String getColumnsStringForSelect( + SpannerPersistentEntity spannerPersistentEntity, + SpannerMappingContext mappingContext, + boolean fetchInterleaved, + boolean forUpdate) { final String sql = String.join(", ", spannerPersistentEntity.columns()); return fetchInterleaved - ? sql + getChildrenSubquery(spannerPersistentEntity, mappingContext) + ? sql + getChildrenSubquery(spannerPersistentEntity, mappingContext, forUpdate) : sql; } @@ -519,7 +616,9 @@ private static String getWhere( } private static String getChildrenSubquery( - SpannerPersistentEntity spannerPersistentEntity, SpannerMappingContext mappingContext) { + SpannerPersistentEntity spannerPersistentEntity, + SpannerMappingContext mappingContext, + boolean forUpdate) { StringJoiner joiner = new StringJoiner(", ", ", ", "").setEmptyValue(""); spannerPersistentEntity.doWithInterleavedProperties( spannerPersistentProperty -> { @@ -533,7 +632,8 @@ private static String getChildrenSubquery( spannerPersistentEntity, mappingContext, spannerPersistentProperty.getColumnName(), - getWhere(spannerPersistentProperty, childPersistentEntity))); + getWhere(spannerPersistentProperty, childPersistentEntity), + forUpdate)); } }); return joiner.toString(); @@ -543,14 +643,15 @@ private static SqlStringAndPlaceholders buildPartTreeSqlString( PartTree tree, SpannerMappingContext spannerMappingContext, Class type, - ParameterAccessor params) { + ParameterAccessor params, + boolean forUpdate) { SpannerPersistentEntity persistentEntity = spannerMappingContext.getPersistentEntityOrFail(type); List tags = new ArrayList<>(); StringBuilder stringBuilder = new StringBuilder(); - buildSelect(persistentEntity, tree, stringBuilder, spannerMappingContext); + buildSelect(persistentEntity, tree, stringBuilder, spannerMappingContext, forUpdate); buildFrom(persistentEntity, stringBuilder); buildWhere(tree, persistentEntity, tags, stringBuilder); applySort( @@ -558,6 +659,9 @@ private static SqlStringAndPlaceholders buildPartTreeSqlString( stringBuilder, persistentEntity); buildLimit(tree, stringBuilder, params.getPageable()); + if (!(tree.isCountProjection() || tree.isExistsProjection())) { + stringBuilder.append(forUpdate ? " FOR UPDATE" : ""); + } String selectSql = stringBuilder.toString(); @@ -575,7 +679,8 @@ private static void buildSelect( SpannerPersistentEntity spannerPersistentEntity, PartTree tree, StringBuilder stringBuilder, - SpannerMappingContext mappingContext) { + SpannerMappingContext mappingContext, + boolean forUpdate) { stringBuilder .append("SELECT ") .append(tree.isDistinct() ? "DISTINCT " : "") @@ -583,7 +688,8 @@ private static void buildSelect( getColumnsStringForSelect( spannerPersistentEntity, mappingContext, - !(tree.isExistsProjection() || tree.isCountProjection()))) + !(tree.isExistsProjection() || tree.isCountProjection()), + forUpdate)) .append(" "); } diff --git a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SqlSpannerQuery.java b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SqlSpannerQuery.java index ce7ae2ed2b9..e5a3dfce951 100644 --- a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SqlSpannerQuery.java +++ b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/query/SqlSpannerQuery.java @@ -196,6 +196,9 @@ private void resolveSpelTags(QueryTagValue queryTagValue) { @Override public List executeRawResult(Object[] parameters) { + if (this.isDml && this.queryMethod.isForUpdate()) { + throw new SpannerDataException("FOR UPDATE cannot be used with a DML query."); + } ParameterAccessor paramAccessor = new ParametersParameterAccessor(getQueryMethod().getParameters(), parameters); @@ -219,6 +222,7 @@ public List executeRawResult(Object[] parameters) { private List executeReadSql(Pageable pageable, Sort sort, QueryTagValue queryTagValue) { SpannerPageableQueryOptions spannerQueryOptions = new SpannerPageableQueryOptions().setAllowPartialRead(true); + spannerQueryOptions.setForUpdate(this.queryMethod.isForUpdate()); if (sort != null && sort.isSorted()) { spannerQueryOptions.setSort(sort); @@ -239,6 +243,9 @@ private List executeReadSql(Pageable pageable, Sort sort, QueryTagValue queryTag queryTagValue.sql, this.spannerMappingContext, entity != null && entity.hasEagerlyLoadedProperties()); + queryTagValue.sql = + SpannerStatementQueryExecutor.applyForUpdate( + queryTagValue.sql, spannerQueryOptions.isForUpdate()); Statement statement = buildStatementFromQueryAndTags(queryTagValue); diff --git a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/support/SimpleSpannerRepository.java b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/support/SimpleSpannerRepository.java index c6396428fb1..b691bcf8306 100644 --- a/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/support/SimpleSpannerRepository.java +++ b/spring-cloud-gcp-data-spanner/src/main/java/com/google/cloud/spring/data/spanner/repository/support/SimpleSpannerRepository.java @@ -20,8 +20,11 @@ import com.google.cloud.spanner.KeySet; import com.google.cloud.spring.data.spanner.core.SpannerOperations; import com.google.cloud.spring.data.spanner.core.SpannerPageableQueryOptions; +import com.google.cloud.spring.data.spanner.core.SpannerQueryOptions; import com.google.cloud.spring.data.spanner.core.SpannerTemplate; +import com.google.cloud.spring.data.spanner.core.mapping.SpannerPersistentEntity; import com.google.cloud.spring.data.spanner.repository.SpannerRepository; +import com.google.cloud.spring.data.spanner.repository.query.SpannerStatementQueryExecutor; import java.util.Collections; import java.util.Optional; import java.util.function.Function; @@ -96,6 +99,27 @@ public Optional findById(I id) { return Optional.ofNullable(result); } + @Override + public Optional findByIdForUpdate(I id) { + Assert.notNull(id, NON_NULL_ID_REQUIRED); + SpannerPersistentEntity persistentEntity = + this.spannerTemplate.getMappingContext().getPersistentEntityOrFail(this.entityType); + return this.spannerTemplate + .query( + this.entityType, + SpannerStatementQueryExecutor.buildQuery( + KeySet.singleKey(toKey(id)), + persistentEntity, + this.spannerTemplate.getSpannerEntityProcessor().getWriteConverter(), + this.spannerTemplate.getMappingContext(), + persistentEntity.getWhere(), + null, + true), + new SpannerQueryOptions().setForUpdate(true)) + .stream() + .findFirst(); + } + @Override public boolean existsById(I id) { Assert.notNull(id, NON_NULL_ID_REQUIRED); diff --git a/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/core/SpannerSortPageQueryOptionsTests.java b/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/core/SpannerSortPageQueryOptionsTests.java index c08442bedd1..4aa7b36a8d6 100644 --- a/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/core/SpannerSortPageQueryOptionsTests.java +++ b/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/core/SpannerSortPageQueryOptionsTests.java @@ -55,4 +55,13 @@ void addQueryOptionTest() { spannerQueryOptions.addQueryOption(r1).addQueryOption(r2); assertThat(Arrays.asList(spannerQueryOptions.getOptions())).containsExactlyInAnyOrder(r1, r2); } + + @Test + void forUpdateTest() { + SpannerQueryOptions spannerQueryOptions = new SpannerQueryOptions(); + + assertThat(spannerQueryOptions.isForUpdate()).isFalse(); + assertThat(spannerQueryOptions.setForUpdate(true)).isSameAs(spannerQueryOptions); + assertThat(spannerQueryOptions.isForUpdate()).isTrue(); + } } diff --git a/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/core/SpannerTemplateTests.java b/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/core/SpannerTemplateTests.java index f4ae511d923..afc5dd68021 100644 --- a/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/core/SpannerTemplateTests.java +++ b/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/core/SpannerTemplateTests.java @@ -56,6 +56,7 @@ import com.google.cloud.spring.data.spanner.core.mapping.Embedded; import com.google.cloud.spring.data.spanner.core.mapping.Interleaved; import com.google.cloud.spring.data.spanner.core.mapping.PrimaryKey; +import com.google.cloud.spring.data.spanner.core.mapping.SpannerDataException; import com.google.cloud.spring.data.spanner.core.mapping.SpannerMappingContext; import com.google.cloud.spring.data.spanner.core.mapping.Table; import com.google.cloud.spring.data.spanner.core.mapping.Where; @@ -513,6 +514,18 @@ void queryTest() { verify(this.databaseClient, times(1)).singleUse(); } + @Test + void forUpdateQueryRequiresReadWriteTransactionTest() { + assertThatThrownBy( + () -> + this.spannerTemplate.query( + TestEntity.class, + Statement.of("SELECT * FROM custom_test_table FOR UPDATE"), + new SpannerQueryOptions().setForUpdate(true))) + .isInstanceOf(SpannerDataException.class) + .hasMessage("FOR UPDATE queries must be executed in a read-write transaction."); + } + @Test void queryFuncTest() { ResultSet resultSet = mock(ResultSet.class); diff --git a/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/query/SpannerQueryLookupStrategyTests.java b/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/query/SpannerQueryLookupStrategyTests.java index fd3a379756d..473f746c973 100644 --- a/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/query/SpannerQueryLookupStrategyTests.java +++ b/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/query/SpannerQueryLookupStrategyTests.java @@ -93,6 +93,11 @@ public String value() { public boolean dmlStatement() { return false; } + + @Override + public boolean forUpdate() { + return false; + } }); } diff --git a/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/query/SpannerStatementQueryTests.java b/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/query/SpannerStatementQueryTests.java index 04013401071..9cb5e42ea92 100644 --- a/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/query/SpannerStatementQueryTests.java +++ b/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/query/SpannerStatementQueryTests.java @@ -33,6 +33,7 @@ import com.google.cloud.spring.data.spanner.core.convert.SpannerEntityProcessor; import com.google.cloud.spring.data.spanner.core.convert.SpannerWriteConverter; import com.google.cloud.spring.data.spanner.core.mapping.Column; +import com.google.cloud.spring.data.spanner.core.mapping.Interleaved; import com.google.cloud.spring.data.spanner.core.mapping.PrimaryKey; import com.google.cloud.spring.data.spanner.core.mapping.SpannerMappingContext; import com.google.cloud.spring.data.spanner.core.mapping.Table; @@ -293,6 +294,52 @@ void sortTest() throws NoSuchMethodException { runPageableOrSortTest(params, method, expectedSql); } + @Test + void forUpdateDerivedQueryTest() throws NoSuchMethodException { + when(this.queryMethod.getName()).thenReturn("findById"); + when(this.queryMethod.isForUpdate()).thenReturn(true); + Method method = QueryHolder.class.getMethod("repositoryMethod10", String.class); + when(this.queryMethod.getQueryMethod()).thenReturn(method); + doReturn(new DefaultParameters(ParametersSource.of(method))) + .when(this.queryMethod) + .getParameters(); + this.partTreeSpannerQuery = spy(createQuery()); + + when(this.spannerTemplate.query((Class) any(), any(), any())) + .thenAnswer( + invocation -> { + Statement statement = invocation.getArgument(1); + assertThat(statement.getSql()) + .isEqualTo( + "SELECT shares, trader_id, ticker, price, action, id, value, uuid " + + "FROM trades WHERE ( id=@tag0 ) FOR UPDATE"); + return Collections.emptyList(); + }); + + doReturn(Object.class).when(this.partTreeSpannerQuery).getReturnedSimpleConvertableItemType(); + doReturn(null).when(this.partTreeSpannerQuery).convertToSimpleReturnType(any(), any()); + + this.partTreeSpannerQuery.execute(new Object[] {"id"}); + } + + @Test + void forUpdateLocksEagerInterleavedChildrenTest() { + SpannerMappingContext mappingContext = new SpannerMappingContext(); + Statement statement = + SpannerStatementQueryExecutor.buildQuery( + com.google.cloud.spanner.KeySet.singleKey(com.google.cloud.spanner.Key.of("parent")), + mappingContext.getPersistentEntityOrFail(Parent.class), + new SpannerWriteConverter(), + mappingContext, + null, + null, + true); + + assertThat(statement.getSql()) + .contains("FROM children WHERE children.id = parents.id FOR UPDATE") + .endsWith("WHERE (id = @tag0) FOR UPDATE"); + } + @Test void uuidUntypedBindingTest() throws NoSuchMethodException { when(this.queryMethod.getName()).thenReturn("findByUuid"); @@ -552,6 +599,22 @@ enum Action { java.util.UUID uuid; } + @Table(name = "parents") + private static class Parent { + @PrimaryKey String id; + + @Interleaved List children; + } + + @Table(name = "children") + private static class Child { + @PrimaryKey(keyOrder = 1) + String id; + + @PrimaryKey(keyOrder = 2) + String childId; + } + // The methods in this class are used to emulate repository methods private static class QueryHolder { public long repositoryMethod1( @@ -617,5 +680,9 @@ public long repositoryMethod8(java.util.UUID tag0) { public long repositoryMethod9(List tag0) { return 0; } + + public long repositoryMethod10(String tag0) { + return 0; + } } } diff --git a/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/support/SimpleSpannerRepositoryTests.java b/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/support/SimpleSpannerRepositoryTests.java index 28beef7d94e..ce1c702a850 100644 --- a/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/support/SimpleSpannerRepositoryTests.java +++ b/spring-cloud-gcp-data-spanner/src/test/java/com/google/cloud/spring/data/spanner/repository/support/SimpleSpannerRepositoryTests.java @@ -102,6 +102,16 @@ void findNullIdTest() { .hasMessage("A non-null ID is required."); } + @Test + void findForUpdateNullIdTest() { + SimpleSpannerRepository spannerRepository = + new SimpleSpannerRepository(this.template, Object.class); + + assertThatThrownBy(() -> spannerRepository.findByIdForUpdate(null)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("A non-null ID is required."); + } + @Test void existsNullIdTest() {