diff --git a/document-store/src/main/java/org/hypertrace/core/documentstore/postgres/FlatPostgresCollection.java b/document-store/src/main/java/org/hypertrace/core/documentstore/postgres/FlatPostgresCollection.java index cbbcbfa3..2c8f263c 100644 --- a/document-store/src/main/java/org/hypertrace/core/documentstore/postgres/FlatPostgresCollection.java +++ b/document-store/src/main/java/org/hypertrace/core/documentstore/postgres/FlatPostgresCollection.java @@ -28,6 +28,7 @@ import java.util.ArrayList; import java.util.Arrays; import java.util.Collection; +import java.util.Comparator; import java.util.HashMap; import java.util.HashSet; import java.util.Iterator; @@ -37,6 +38,7 @@ import java.util.Map.Entry; import java.util.Optional; import java.util.Set; +import java.util.TreeMap; import java.util.stream.Collectors; import org.hypertrace.core.documentstore.BulkArrayValueUpdateRequest; import org.hypertrace.core.documentstore.BulkDeleteResult; @@ -360,8 +362,10 @@ public boolean bulkUpsert(Map documents) { PostgresDataType pkType = getPrimaryKeyType(tableName, pkColumn); try { - // Parse all documents - Map parsedDocuments = new LinkedHashMap<>(); + // TreeMap keyed by Key.toString() so the downstream JDBC batch iterates rows in a canonical + // order. Postgres acquires row locks in batch-entry order, so overlapping concurrent batches + // must share an ordering to avoid deadlocks. + Map parsedDocuments = new TreeMap<>(Comparator.comparing(Key::toString)); List ignoredDocuments = new ArrayList<>(); for (Map.Entry entry : documents.entrySet()) { @@ -454,8 +458,10 @@ public boolean bulkCreateOrReplace(Map documents) { PostgresDataType pkType = getPrimaryKeyType(tableName, pkColumn); try { - // Parse all documents - Map parsedDocuments = new LinkedHashMap<>(); + // TreeMap keyed by Key.toString() so the downstream JDBC batch iterates rows in a canonical + // order. Postgres acquires row locks in batch-entry order, so overlapping concurrent batches + // must share an ordering to avoid deadlocks. + Map parsedDocuments = new TreeMap<>(Comparator.comparing(Key::toString)); List ignoredDocuments = new ArrayList<>(); for (Map.Entry entry : documents.entrySet()) { @@ -1022,6 +1028,13 @@ private int executeBatchUpdate( List keys = keyGroup.getKeys(); List> allKeyParams = keyGroup.getKeyParams(); + Integer[] order = new Integer[keys.size()]; + for (int i = 0; i < order.length; i++) { + order[i] = i; + } + // Sort keys to ensure consistent order to avoid deadlocks + Arrays.sort(order, Comparator.comparing(i -> keys.get(i).toString())); + List setFragments = new ArrayList<>(keyGroup.getSetFragments()); List timestampParam = new ArrayList<>(); appendLastUpdatedTimestamp(setFragments, timestampParam, tableName, epochMillis); @@ -1034,7 +1047,7 @@ private int executeBatchUpdate( LOGGER.debug("Executing batch update SQL: {} for {} keys", sql, keys.size()); try (PreparedStatement ps = connection.prepareStatement(sql)) { - for (int i = 0; i < keys.size(); i++) { + for (int i : order) { int idx = 1; for (Object param : allKeyParams.get(i)) { ps.setObject(idx++, param);