From 7a071a62289ef61727dc25f9dd84b9400d5b1468 Mon Sep 17 00:00:00 2001 From: godotg Date: Fri, 7 Oct 2022 21:35:48 +0800 Subject: [PATCH] feat[hashmap]: support primitive hash map --- .../protocol/collection/HashMapIntInt.java | 428 ++++++++++++++++++ .../zfoo/protocol/collection/HashMapTest.java | 327 +++++++++++++ 2 files changed, 755 insertions(+) create mode 100644 protocol/src/main/java/com/zfoo/protocol/collection/HashMapIntInt.java create mode 100644 protocol/src/test/java/com/zfoo/protocol/collection/HashMapTest.java diff --git a/protocol/src/main/java/com/zfoo/protocol/collection/HashMapIntInt.java b/protocol/src/main/java/com/zfoo/protocol/collection/HashMapIntInt.java new file mode 100644 index 00000000..55f05f3e --- /dev/null +++ b/protocol/src/main/java/com/zfoo/protocol/collection/HashMapIntInt.java @@ -0,0 +1,428 @@ +/* + * Copyright (C) 2020 The zfoo Authors + * Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except + * in compliance with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software distributed under the License is distributed + * on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and limitations under the License. + */ + +package com.zfoo.protocol.collection; + +import com.zfoo.protocol.util.StringUtils; +import io.netty.util.collection.IntObjectHashMap; +import io.netty.util.internal.MathUtil; + +import java.util.*; + + +/** + * @author godotg + * @version 3.0 + */ +public class HashMapIntInt implements Map { + + public static final byte FREE = 0; + public static final byte REMOVED = 1; + public static final byte FILLED = 2; + + private int[] keys; + private int[] values; + private byte[] statuses; + private int size; + private int maxSize; + private int mask; + + /** + * Calculates the maximum size allowed before rehashing. + */ + public static int calcMaxSize(int capacity) { + // Clip the upper bound so that there will always be at least one available slot. + int upperBound = capacity - 1; + return Math.min(upperBound, (int) (capacity * IntObjectHashMap.DEFAULT_LOAD_FACTOR)); + } + + /** + * Get the next sequential index after index and wraps if necessary. + */ + public static int probeNext(int index, int mask) { + // The array lengths are always a power of two, so we can use a bitmask to stay inside the array bound + return (index + 1) & mask; + } + + public HashMapIntInt() { + this(IntObjectHashMap.DEFAULT_CAPACITY); + } + + public HashMapIntInt(int initialCapacity) { + var capacity = MathUtil.safeFindNextPositivePowerOfTwo(initialCapacity); + initCapacity(capacity); + } + + private void initCapacity(int capacity) { + mask = capacity - 1; + + keys = new int[capacity]; + values = new int[capacity]; + statuses = new byte[capacity]; + + maxSize = calcMaxSize(capacity); + } + + private void ensureCapacity() { + if (size > maxSize) { + if (keys.length == Integer.MAX_VALUE) { + throw new IllegalStateException("Max capacity reached at size=" + size); + } + // Double the capacity. + rehash(keys.length << 1); + } + } + + + @Override + public int size() { + return size; + } + + @Override + public boolean isEmpty() { + return size == 0; + } + + public boolean containsKeyPrimitive(int key) { + return indexOf(key) >= 0; + } + + @Override + public boolean containsKey(Object key) { + return containsKeyPrimitive(ArrayUtils.intValue((Integer) key)); + } + + @Override + public boolean containsValue(Object value) { + return containsValuePrimitive(ArrayUtils.intValue((Integer) value)); + } + + public boolean containsValuePrimitive(int value) { + for (var i = 0; i < statuses.length; i++) { + if (statuses[i] == FILLED && values[i] == value) { + return true; + } + } + return false; + } + + @Override + public Integer get(Object key) { + var index = indexOf(ArrayUtils.intValue((Integer) key)); + return index == -1 ? null : values[index]; + } + + @Override + public Integer put(Integer key, Integer value) { + return putPrimitive(ArrayUtils.intValue(key), ArrayUtils.intValue(value)); + } + + public Integer putPrimitive(int key, int value) { + var startIndex = hashIndex(key); + var index = startIndex; + + var firstRemoveIndex = -1; + for (; ; ) { + var status = statuses[index]; + if (status == FREE) { + index = firstRemoveIndex < 0 ? index : firstRemoveIndex; + set(index, key, value, FILLED); + size++; + ensureCapacity(); + return null; + } else if (status == REMOVED) { + firstRemoveIndex = firstRemoveIndex < 0 ? index : firstRemoveIndex; + } else if (keys[index] == key) { // status == FILLED + // Found existing entry with this key, just replace the value. + var previousValue = values[index]; + values[index] = value; + return previousValue; + } + + // Conflict, keep probing ... + if ((index = probeNext(index, mask)) == startIndex) { + if (firstRemoveIndex < 0) { + throw new IllegalStateException("Unable to insert, the map was full at MAX_ARRAY_SIZE and couldn't grow"); + } else { + set(firstRemoveIndex, key, value, FILLED); + size++; + ensureCapacity(); + return null; + } + } + } + } + + @Override + public Integer remove(Object key) { + return removePrimitive(ArrayUtils.intValue((Integer) key)); + } + + public Integer removePrimitive(int key) { + var index = indexOf(key); + if (index == -1) { + return null; + } + var prev = values[index]; + removeAt(index); + return prev; + } + + private void removeAt(int index) { + set(index, 0, 0, REMOVED); + size--; + } + + @Override + public void putAll(Map m) { + for (Entry entry : m.entrySet()) { + put(entry.getKey(), entry.getValue()); + } + } + + @Override + public void clear() { + Arrays.fill(keys, 0); + Arrays.fill(values, 0); + Arrays.fill(statuses, FREE); + size = 0; + } + + @Override + public Set keySet() { + return new KeySet(); + } + + @Override + public Collection values() { + return new ValueSet(); + } + + @Override + public Set> entrySet() { + return new EntrySet(); + } + + private int hashIndex(int key) { + return key & mask; + } + + private void set(int index, int key, int value, byte status) { + keys[index] = key; + values[index] = value; + statuses[index] = status; + } + + private void rehash(int newCapacity) { + var oldKeys = keys; + var oldValues = values; + var oldStatuses = statuses; + + initCapacity(newCapacity); + + for (var i = 0; i < oldStatuses.length; ++i) { + var oldStatus = oldStatuses[i]; + if (oldStatus == FILLED) { + var oldKey = oldKeys[i]; + var oldValue = oldValues[i]; + int index = hashIndex(oldKey); + + for (; ; ) { + if (statuses[index] == FREE) { + set(index, oldKey, oldValue, FILLED); + break; + } + + index = probeNext(index, mask); + } + } + } + } + + private int indexOf(int key) { + int startIndex = hashIndex(key); + int index = startIndex; + + for (; ; ) { + var status = statuses[index]; + if (status == FREE) { + // It's available, so no chance that this value exists anywhere in the map. + return -1; + } + if (key == keys[index] && status == FILLED) { + return index; + } + + // Conflict, keep probing ... + if ((index = probeNext(index, mask)) == startIndex) { + return -1; + } + } + } + + private class PrimitiveEntry implements Entry { + int entryIndex; + + PrimitiveEntry(int entryIndex) { + this.entryIndex = entryIndex; + } + + @Override + public Integer getKey() { + return keys[entryIndex]; + } + + @Override + public Integer getValue() { + return values[entryIndex]; + } + + @Override + public Integer setValue(Integer value) { + var prevValue = values[entryIndex]; + values[entryIndex] = value; + return prevValue; + } + } + + private class FastIterator implements Iterator> { + int lastCursor = -1; + int cursor = -1; + + private void scanNext() { + while (++cursor != statuses.length && statuses[cursor] != FILLED) { + } + } + + @Override + public boolean hasNext() { + if (cursor == -1) { + scanNext(); + } + return cursor != statuses.length; + } + + @Override + public Entry next() { + if (!hasNext()) { + throw new NoSuchElementException(); + } + + lastCursor = cursor; + scanNext(); + + return new PrimitiveEntry(lastCursor); + } + + @Override + public void remove() { + if (lastCursor == -1) { + throw new IllegalStateException("next must be called before each remove."); + } + removeAt(lastCursor); + cursor = -1; + lastCursor = -1; + } + } + + private final class KeySet extends AbstractSet { + FastIterator fastIterator = new FastIterator(); + + @Override + public Iterator iterator() { + return new Iterator() { + @Override + public boolean hasNext() { + return fastIterator.hasNext(); + } + + @Override + public Integer next() { + return fastIterator.next().getKey(); + } + + @Override + public void remove() { + fastIterator.remove(); + } + }; + } + + @Override + public int size() { + return HashMapIntInt.this.size(); + } + } + + private final class ValueSet extends AbstractSet { + FastIterator fastIterator = new FastIterator(); + + @Override + public Iterator iterator() { + return new Iterator() { + @Override + public boolean hasNext() { + return fastIterator.hasNext(); + } + + @Override + public Integer next() { + return fastIterator.next().getValue(); + } + + @Override + public void remove() { + fastIterator.remove(); + } + }; + } + + @Override + public int size() { + return HashMapIntInt.this.size(); + } + } + + private final class EntrySet extends AbstractSet> { + @Override + public Iterator> iterator() { + return new FastIterator(); + } + + @Override + public int size() { + return HashMapIntInt.this.size(); + } + } + + @Override + public String toString() { + if (isEmpty()) { + return StringUtils.EMPTY_JSON; + } + var builder = new StringBuilder(4 * size); + builder.append('{'); + var first = true; + for (int i = 0; i < values.length; ++i) { + if (statuses[i] != FILLED) { + continue; + } + if (!first) { + builder.append(", "); + } + builder.append(keys[i]).append('=').append(values[i]); + first = false; + } + return builder.append('}').toString(); + } +} diff --git a/protocol/src/test/java/com/zfoo/protocol/collection/HashMapTest.java b/protocol/src/test/java/com/zfoo/protocol/collection/HashMapTest.java new file mode 100644 index 00000000..cc3e52fe --- /dev/null +++ b/protocol/src/test/java/com/zfoo/protocol/collection/HashMapTest.java @@ -0,0 +1,327 @@ +/* + * Copyright (C) 2020 The zfoo Authors + * Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except + * in compliance with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software distributed under the License is distributed + * on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and limitations under the License. + */ + +package com.zfoo.protocol.collection; + +import org.junit.Assert; +import org.junit.Test; + +import java.util.*; + +/** + * @author godotg + * @version 3.0 + */ +public class HashMapTest { + + @Test + public void putTest() { + var map = new HashMapIntInt(0); + map.put(0, 0); + map.put(1, 1); + map.put(2, 2); + map.put(3, 3); + map.remove(0); + map.remove(1); + map.remove(2); + map.remove(3); + map.put(0, 0); + map.put(1, 1); + map.put(2, 2); + map.put(3, 3); + System.out.println(map.size()); + System.out.println(map); + } + + private void assertKey(HashMapIntInt primitiveMap, HashMap javaMap, int key) { + var primitiveValue = primitiveMap.get(key); + var javaValue = javaMap.get(key); + Assert.assertEquals(primitiveValue, javaValue); + Assert.assertEquals(primitiveMap.containsKey(key), javaMap.containsKey(key)); + Assert.assertEquals(primitiveMap.size(), javaMap.size()); + } + + @Test + public void testPutGetRemoveContainsSmallMaps() { + var random = new Random(322); + for (var it = 0; it < 10000; it++) { + var primitiveMap = new HashMapIntInt(0); + var javaMap = new HashMap(); + Assert.assertTrue(primitiveMap.isEmpty()); + for (var i = 0; i < 100; i++) { + var key = random.nextInt(50); + var value = random.nextInt(50); + assertKey(primitiveMap, javaMap, key); + if (random.nextBoolean()) { + var primitiveValue = primitiveMap.put(key, value); + var javaValue = javaMap.put(key, value); + Assert.assertEquals(primitiveValue, javaValue); + } else { + var primitiveValue = primitiveMap.remove(key); + var javaValue = javaMap.remove(key); + Assert.assertEquals(primitiveValue, javaValue); + } + assertKey(primitiveMap, javaMap, key); + } + } + } + + @Test + public void testPutGetRemoveContainsBigMaps() { + var random = new Random(Integer.MAX_VALUE); + for (var it = 0; it < 100; it++) { + var primitiveMap = new HashMapIntInt(0); + var javaMap = new HashMap(); + Assert.assertTrue(primitiveMap.isEmpty()); + for (var i = 0; i < 10000; i++) { + var key = random.nextInt(Integer.MAX_VALUE); + var value = random.nextInt(Integer.MAX_VALUE); + assertKey(primitiveMap, javaMap, key); + if (random.nextBoolean()) { + var primitiveValue = primitiveMap.put(key, value); + var javaValue = javaMap.put(key, value); + Assert.assertEquals(primitiveValue, javaValue); + } else { + var primitiveValue = primitiveMap.remove(key); + var javaValue = javaMap.remove(key); + Assert.assertEquals(primitiveValue, javaValue); + } + assertKey(primitiveMap, javaMap, key); + } + } + } + + @Test + public void testPut() { + var map = new HashMapIntInt(); + var array = new int[10000]; + Arrays.fill(array, -1); + var random = new Random(32232); + int size = 0; + for (var i = 0; i < 1000000; i++) { + Assert.assertEquals(map.size(), size); + var key = random.nextInt(10000); + Integer oldValue = map.put(key, i); + if (array[key] == -1) { + Assert.assertNull(oldValue); + size++; + } else { + Assert.assertNotNull(oldValue); + Assert.assertEquals(oldValue.intValue(), array[key]); + } + array[key] = i; + Assert.assertEquals(map.size(), size); + } + } + + @Test + public void testGetContainsKey() { + var map = new HashMapIntInt(); + var random = new Random(322322); + int size = 100000; + int[] array = new int[size]; + for (int i = 0; i < size; i++) { + array[i] = random.nextInt(); + var oldValue = map.put(i, array[i]); + Assert.assertNull(oldValue); + } + Assert.assertEquals(map.size(), size); + for (int i = 0; i < size; i++) { + Assert.assertEquals(map.get(i).intValue(), array[i]); + Assert.assertTrue(map.containsKey(i)); + map.get(~i); + Assert.assertFalse(map.containsKey(~i)); + } + } + + @Test + public void testRemove() { + var map = new HashMapIntInt(); + var random = new Random(3223223); + int size = 100000; + int[] array = new int[size]; + for (int i = 0; i < size; i++) { + array[i] = random.nextInt(); + var oldValue = map.put(i, array[i]); + Assert.assertNull(oldValue); + } + Assert.assertEquals(map.size(), size); + for (int i = 0; i < size; i++) { + Assert.assertTrue(map.containsKey(i)); + Assert.assertEquals(map.remove(i).intValue(), array[i]); + Assert.assertFalse(map.containsKey(i)); + map.remove(i); + Assert.assertFalse(map.containsKey(i)); + } + } + + @Test + public void testClear() { + var map = new HashMapIntInt(); + Assert.assertTrue(map.isEmpty()); + Assert.assertEquals(map.size(), 0); + int size = 0; + for (int i = 1; i <= 1000000; i++) { + map.put(i, -i); + size++; + Assert.assertFalse(map.isEmpty()); + Assert.assertEquals(map.size(), size); + if ((i & (i - 1)) == 0) { + map.clear(); + Assert.assertTrue(map.isEmpty()); + Assert.assertEquals(map.size(), 0); + size = 0; + } + } + } + + @Test + public void testKeysValuesArrays() { + var random = new Random(32232232); + var n = 100000; + int[] keys = new int[n]; + int[] values = new int[n]; + for (int i = 0; i < n; i++) { + keys[i] = (1 + i) * 10000 + random.nextInt(9000); + values[i] = random.nextInt(); + } + var map = new HashMapIntInt(); + for (int i = 0; i < n; i++) { + var oldValue = map.put(keys[i], values[i]); + Assert.assertNull(oldValue); + } + Assert.assertEquals(map.size(), n); + var mapKeys = new ArrayList<>(map.keySet()); + var mapValues = new ArrayList<>(map.values()); + Collections.sort(mapKeys); + Collections.sort(mapValues); + Arrays.sort(keys); + Arrays.sort(values); + Assert.assertArrayEquals(ArrayUtils.intToArray(mapKeys), keys); + Assert.assertArrayEquals(ArrayUtils.intToArray(mapValues), values); + } + + @Test + public void testConstructors() { + var srcMap = new HashMapIntInt(); + var javaMap = new HashMap(); + int n = 10000; + for (int i = 0; i < n; i++) { + srcMap.put(~i, i); + javaMap.put(~i, i); + } + var map1 = new HashMapIntInt(); + map1.putAll(srcMap); + var map2 = new HashMapIntInt(); + map2.putAll(javaMap); + Assert.assertEquals(map1.size(), n); + Assert.assertEquals(map2.size(), n); + for (int i = 0; i < n; i++) { + Assert.assertEquals(map1.get(~i).intValue(), i); + Assert.assertEquals(map2.get(~i).intValue(), i); + } + Assert.assertEquals(map1.size(), srcMap.size()); + Assert.assertEquals(map2.size(), srcMap.size()); + } + + @Test + public void testIterator() { + var random = new Random(322322322); + int n = 1000; + int[] array = new int[n]; + for (int i = 0; i < n; i++) { + array[i] = random.nextInt(); + } + var map = new HashMapIntInt(); + for (int i = 0; i < n; i++) { + var oldValue = map.put(i, array[i]); + Assert.assertNull(oldValue); + } + boolean[] visited = new boolean[n]; + var it = map.entrySet().iterator(); + for (int i = 0; i < n; i++) { + Assert.assertTrue(it.hasNext()); + var next = it.next(); + int key = next.getKey(); + int value = next.getValue(); + Assert.assertEquals(value, array[key]); + Assert.assertFalse(visited[key]); + visited[key] = true; + for (int retries = 0; retries < 2; retries++) { + Assert.assertEquals(next.getKey().intValue(), key); + Assert.assertEquals(next.getValue().intValue(), value); + } + } + for (int retries = 0; retries < 3; retries++) { + Assert.assertFalse(it.hasNext()); + try { + it.next().getKey(); + } catch (NoSuchElementException e) { + // as expected + } + try { + it.next().getValue(); + } catch (NoSuchElementException e) { + // as expected + } + } + + it = map.entrySet().iterator(); + while (it.hasNext()) { + var entry = it.next(); + it.remove(); + Assert.assertFalse(map.containsKey(entry.getKey())); + } + Assert.assertTrue(map.isEmpty()); + } + + + @Test + public void testCompressingAfterRemoving() { + int n = 1000000; + int[] a = new int[n]; + for (int i = 0; i < n; i++) { + a[i] = i; + } + Random rnd = new Random(3223223223L); + for (int i = 0; i < n; i++) { + int j = i + rnd.nextInt(n - i); + int tmp = a[i]; + a[i] = a[j]; + a[j] = tmp; + } + var map = new HashMapIntInt(); + for (int i = 0; i < n; i++) { + map.put(a[i], i); + } + for (int i = 0; i < n - 1000; i++) { + var oldValue = map.remove(a[i]); + Assert.assertNotNull(oldValue); + } + Assert.assertEquals(map.size(), 1000); + + Assert.assertEquals(map.size(), 1000); + // Length of the arrays in the map must be O(size). If it's not, there will be a timeout + int dummy1 = 0, dummy2 = 0; + for (int i = 0; i < 1000; i++) { + var it = map.entrySet().iterator(); + while(it.hasNext()) { + var next = it.next(); + dummy1 ^= next.getKey(); + dummy2 ^= next.getValue(); + } + } + Assert.assertEquals(dummy1, 0); + Assert.assertEquals(dummy2, 0); + } + +}