From 08d07b0d9d5db780f41ab783f86389f329484948 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Tue, 4 Jun 2024 10:08:04 +0300 Subject: more on unit testing and some refactoring --- internal/server/handler/handler.go | 2 +- internal/server/repository/repository.go | 34 +++++++++----- internal/server/repository/repository_test.go | 64 +++++++++++++++++++++++++++ internal/types/entry.go | 26 ++++++----- internal/types/entry_test.go | 6 +-- 5 files changed, 106 insertions(+), 26 deletions(-) create mode 100644 internal/server/repository/repository_test.go (limited to 'internal') diff --git a/internal/server/handler/handler.go b/internal/server/handler/handler.go index 6549c9f..12c6660 100644 --- a/internal/server/handler/handler.go +++ b/internal/server/handler/handler.go @@ -62,7 +62,7 @@ func (h Handler) Get(w http.ResponseWriter, r *http.Request) error { return fmt.Errorf("invalid id %s", id) } - data, err := repository.Instance(h.conf.DataDir).Get(id) + data, err := repository.Instance(h.conf.DataDir).GetBytes(id) if err != err { return err } diff --git a/internal/server/repository/repository.go b/internal/server/repository/repository.go index b98ccfa..8e84f4b 100644 --- a/internal/server/repository/repository.go +++ b/internal/server/repository/repository.go @@ -11,7 +11,7 @@ import ( ) var ( - instance *Repository + instance Repository once sync.Once ) @@ -33,19 +33,23 @@ type Repository struct { fs fs } -func Instance(dataDir string) *Repository { +func Instance(dataDir string) Repository { once.Do(func() { - instance = &Repository{ - dataDir: dataDir, - entries: make(map[string]types.Entry), - mu: &sync.Mutex{}, - fs: vfs.RealFS{}, - } + instance = newRepository(dataDir, vfs.RealFS{}) }) return instance } -func (r Repository) add(entry types.Entry) { +func newRepository(dataDir string, fs fs) Repository { + return Repository{ + dataDir: dataDir, + entries: make(map[string]types.Entry), + mu: &sync.Mutex{}, + fs: fs, + } +} + +func (r Repository) put(entry types.Entry) { r.mu.Lock() defer r.mu.Unlock() r.entries[entry.ID] = entry @@ -63,7 +67,7 @@ func (r Repository) load() error { if err != err { return err } - r.add(entry) + r.put(entry) } return nil @@ -85,10 +89,18 @@ func (r Repository) List() ([]byte, error) { return json.Marshal(pairs) } -func (r Repository) Get(id string) ([]byte, error) { +func (r Repository) GetBytes(id string) ([]byte, error) { return r.fs.ReadFile(fmt.Sprintf("%s/%s", r.dataDir, id)) } +func (r Repository) Get(id string) (types.Entry, error) { + bytes, err := r.GetBytes(id) + if err != nil { + return types.Entry{}, err + } + return types.NewEntry(bytes) +} + func (r Repository) HasSameEntry(pair EntryPair) bool { r.mu.Lock() defer r.mu.Unlock() diff --git a/internal/server/repository/repository_test.go b/internal/server/repository/repository_test.go new file mode 100644 index 0000000..176808d --- /dev/null +++ b/internal/server/repository/repository_test.go @@ -0,0 +1,64 @@ +package repository + +import ( + "testing" + + "codeberg.org/snonux/gos/internal/types" + "codeberg.org/snonux/gos/internal/vfs" +) + +func TestRepositoryGet(t *testing.T) { + t.Parallel() + + entry, _, err := twoDifferentEntries() + if err != nil { + t.Error(err) + return + } + + repo.put(entry) + t.Log(fs) + + entryGot, err := repo.Get(entry.ID) + if err != nil { + t.Error(err) + return + } + if !entryGot.Equals(entry) { + t.Error("expected to get", entry, "but got", entryGot) + } +} + +// TODO: Write unit tests for the remainder of the repo methods + +func setupRepository() (repo Repository, entry1, entry2 types.Entry, err error) { + fs := make(vfs.MemoryFS) + repo = newRepository("./data", fs) + + entry1Str := ` + { + "Body": "Body text here", + "Shared": [ + { "Name": "Foo", "Is": true }, + { "Name": "Bar", "Is": false } + ] + } + ` + entry1, err = types.NewEntry([]byte(entry1Str), fs) + if err != nil { + return + } + + entry2Str := ` + { + "Body": "Body text here", + "Shared": [ + { "Name": "Foo", "Is": true }, + { "Name": "Bar", "Is": true }, + { "Name": "Baz", "Is": false } + ] + } + ` + entry2, err = types.NewEntry([]byte(entry2Str), fs) + return +} diff --git a/internal/types/entry.go b/internal/types/entry.go index eca7b91..cc2f8dd 100644 --- a/internal/types/entry.go +++ b/internal/types/entry.go @@ -50,36 +50,36 @@ type Entry struct { mu *sync.Mutex } -func NewEntry(bytes []byte) (Entry, error) { +func NewEntry(bytes []byte, fs ...fs) (Entry, error) { var e Entry if err := json.Unmarshal(bytes, &e); err != nil { return e, fmt.Errorf("unable to deserialise payload: %w", err) } - e.initialize() + e.initialize(fs...) if e.ID == "" { e.ID = fmt.Sprintf("%x", sha256.Sum256([]byte(e.Body))) } return e, nil } -func NewEntryFromFile(filePath string, fsToUse ...fs) (Entry, error) { +func NewEntryFromFile(filePath string, fs_ ...fs) (Entry, error) { var ( bytes []byte err error - fs fs = vfs.RealFS{} + fs fs ) - if len(fsToUse) > 0 { - fs = fsToUse[0] + if len(fs_) > 0 { + fs = fs_[0] + } else { + fs = vfs.RealFS{} } bytes, err = fs.ReadFile(filePath) if err != err { return Entry{}, err } - e, err := NewEntry(bytes) - e.fs = fs - return e, err + return NewEntry(bytes, fs) } func NewEntryFromCopy(other Entry) (Entry, error) { @@ -88,10 +88,14 @@ func NewEntryFromCopy(other Entry) (Entry, error) { return e.Update(other) } -func (e *Entry) initialize() { +func (e *Entry) initialize(fs ...fs) { e.mu = &sync.Mutex{} e.checksumDirty = true - e.fs = vfs.RealFS{} + if len(fs) > 1 { + e.fs = fs[0] + } else { + e.fs = vfs.RealFS{} + } } func (e Entry) Equals(other Entry) bool { diff --git a/internal/types/entry_test.go b/internal/types/entry_test.go index c6a7320..3a1deb7 100644 --- a/internal/types/entry_test.go +++ b/internal/types/entry_test.go @@ -21,7 +21,7 @@ func TestEntryChecksum(t *testing.T) { t.Log(entry.Checksum()) } -func twoDifferentEntries(t *testing.T) (entry1, entry2 Entry, err error) { +func twoDifferentEntries() (entry1, entry2 Entry, err error) { entry1Str := ` { "Body": "Body text here", @@ -53,7 +53,7 @@ func twoDifferentEntries(t *testing.T) (entry1, entry2 Entry, err error) { func TestEquals(t *testing.T) { t.Parallel() - entry1, entry2, err := twoDifferentEntries(t) + entry1, entry2, err := twoDifferentEntries() if err != nil { t.Error(err) return @@ -69,7 +69,7 @@ func TestEquals(t *testing.T) { func TestUpdate(t *testing.T) { t.Parallel() - entry1, entry2, err := twoDifferentEntries(t) + entry1, entry2, err := twoDifferentEntries() if err != nil { t.Error(err) return -- cgit v1.2.3