summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
authorPaul Buetow <paul@buetow.org>2024-06-04 10:08:04 +0300
committerPaul Buetow <paul@buetow.org>2024-06-04 10:08:04 +0300
commit08d07b0d9d5db780f41ab783f86389f329484948 (patch)
tree370b67e7f016939f55ae93dca6f379613aeea2bc /internal
parent11788abad75c0d4920ed4f7797febd7d36569a66 (diff)
more on unit testing and some refactoring
Diffstat (limited to 'internal')
-rw-r--r--internal/server/handler/handler.go2
-rw-r--r--internal/server/repository/repository.go34
-rw-r--r--internal/server/repository/repository_test.go64
-rw-r--r--internal/types/entry.go26
-rw-r--r--internal/types/entry_test.go6
5 files changed, 106 insertions, 26 deletions
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