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/types/entry.go | 26 +++++++++++++++----------- internal/types/entry_test.go | 6 +++--- 2 files changed, 18 insertions(+), 14 deletions(-) (limited to 'internal/types') 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