From b632f0be7a1495a6532f73cc152ef14a9c69a315 Mon Sep 17 00:00:00 2001 From: rjl493456442 Date: Tue, 17 Jul 2018 13:27:55 +0800 Subject: [PATCH] ethdb: fix memory database --- ethdb/database_test.go | 22 ++++++++++++++++++++++ ethdb/memory_database.go | 12 ++++++++---- 2 files changed, 30 insertions(+), 4 deletions(-) diff --git a/ethdb/database_test.go b/ethdb/database_test.go index 2deb50988c..50386f41e9 100644 --- a/ethdb/database_test.go +++ b/ethdb/database_test.go @@ -59,6 +59,28 @@ func TestMemoryDB_PutGet(t *testing.T) { func testPutGet(db ethdb.Database, t *testing.T) { t.Parallel() + for _, k := range test_values { + err := db.Put([]byte(k), nil) + if err != nil { + t.Fatalf("put failed: %v", err) + } + } + + for _, k := range test_values { + data, err := db.Get([]byte(k)) + if err != nil { + t.Fatalf("get failed: %v", err) + } + if len(data) != 0 { + t.Fatalf("get returned wrong result, got %q expected nil", string(data)) + } + } + + _, err := db.Get([]byte("non-exist-key")) + if err == nil { + t.Fatalf("expect to return a non found error") + } + for _, v := range test_values { err := db.Put([]byte(v), []byte(v)) if err != nil { diff --git a/ethdb/memory_database.go b/ethdb/memory_database.go index f28ff54818..727f2f7ca3 100644 --- a/ethdb/memory_database.go +++ b/ethdb/memory_database.go @@ -96,7 +96,10 @@ func (db *MemDatabase) NewBatch() Batch { func (db *MemDatabase) Len() int { return len(db.db) } -type kv struct{ k, v []byte } +type kv struct { + k, v []byte + del bool +} type memBatch struct { db *MemDatabase @@ -105,13 +108,14 @@ type memBatch struct { } func (b *memBatch) Put(key, value []byte) error { - b.writes = append(b.writes, kv{common.CopyBytes(key), common.CopyBytes(value)}) + b.writes = append(b.writes, kv{common.CopyBytes(key), common.CopyBytes(value), false}) b.size += len(value) return nil } func (b *memBatch) Delete(key []byte) error { - b.writes = append(b.writes, kv{common.CopyBytes(key), nil}) + b.writes = append(b.writes, kv{common.CopyBytes(key), nil, true}) + b.size += 1 return nil } @@ -120,7 +124,7 @@ func (b *memBatch) Write() error { defer b.db.lock.Unlock() for _, kv := range b.writes { - if kv.v == nil { + if kv.del { delete(b.db.db, string(kv.k)) continue }