Skip to content

Commit efd5c13

Browse files
authored
Merge pull request #13 from jbeshir/incorporate-tests
Tests: Add initial tests, fix missing PINECONE_INDEX_NAME env var
2 parents 0076dac + 1eb8d82 commit efd5c13

26 files changed

Lines changed: 1731 additions & 6 deletions

‎.env.dist‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@ MYSQL_URI=alignment_research_feed:pass@tcp(localhost:3306)/alignment_research_da
1919

2020
SIMILARITY_DRIVER=null
2121
PINECONE_API_KEY=
22+
PINECONE_INDEX_NAME=
2223

2324
AUTH_DRIVER=null
2425
AUTH0_DOMAIN=

‎.mockery.yaml‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,19 @@
1+
with-expecter: true
2+
dir: "{{.InterfaceDir}}/mocks"
3+
mockname: "Mock{{.InterfaceName}}"
4+
outpkg: "mocks"
5+
filename: "mock_{{.InterfaceName | snakecase}}.go"
6+
packages:
7+
github.com/jbeshir/alignment-research-feed/internal/datasources:
8+
interfaces:
9+
ArticleFetcher:
10+
ArticleReadSetter:
11+
ArticleThumbsUpSetter:
12+
ArticleThumbsDownSetter:
13+
LatestArticleLister:
14+
ThumbsUpArticleLister:
15+
UserVectorGetter:
16+
UserVectorSyncer:
17+
SimilarArticleLister:
18+
ArticleVectorFetcher:
19+
SimilarArticlesByVectorLister:

‎Makefile‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,11 +7,16 @@ setup-tools: setup-files
77
go install github.com/sqlc-dev/sqlc/cmd/sqlc@latest
88
go install golang.org/x/tools/cmd/goimports@latest
99
go install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@latest
10+
go install github.com/vektra/mockery/v2@latest
1011

1112
.PHONY generate:
1213
generate:
1314
go generate ./...
1415

16+
.PHONY mocks:
17+
mocks:
18+
mockery
19+
1520
.PHONY test-short:
1621
test-short:
1722
go test -v -short ./...

‎go.mod‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ require (
2424
github.com/huandu/xstrings v1.5.0 // indirect
2525
github.com/oapi-codegen/runtime v1.1.1 // indirect
2626
github.com/pmezard/go-difflib v1.0.0 // indirect
27+
github.com/stretchr/objx v0.5.2 // indirect
2728
golang.org/x/net v0.47.0 // indirect
2829
golang.org/x/sys v0.38.0 // indirect
2930
golang.org/x/text v0.31.0 // indirect

‎go.sum‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,8 @@ github.com/rogpeppe/go-internal v1.9.0 h1:73kH8U+JUqXU8lRuOHeVHaa/SZPifC7BkcraZV
4949
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
5050
github.com/spkg/bom v0.0.0-20160624110644-59b7046e48ad/go.mod h1:qLr4V1qq6nMqFKkMo8ZTx3f+BZEkzsRUY10Xsm2mwU0=
5151
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
52+
github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY=
53+
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
5254
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
5355
github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA=
5456
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=

‎internal/app/app.go‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -69,7 +69,11 @@ func setupSimilarityRepository(ctx context.Context) (datasources.SimilarityRepos
6969
case "null":
7070
return datasources.NullSimilarityRepository{}, nil
7171
case "pinecone":
72-
client, err := pinecone.NewClient(ctx, MustGetEnvAsString(ctx, "PINECONE_API_KEY"))
72+
client, err := pinecone.NewClient(
73+
ctx,
74+
MustGetEnvAsString(ctx, "PINECONE_API_KEY"),
75+
MustGetEnvAsString(ctx, "PINECONE_INDEX_NAME"),
76+
)
7377
if err != nil {
7478
return nil, fmt.Errorf("connecting to pinecone: %w", err)
7579
}
Lines changed: 128 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,128 @@
1+
package command
2+
3+
import (
4+
"context"
5+
"errors"
6+
"log/slog"
7+
"testing"
8+
9+
"github.com/jbeshir/alignment-research-feed/internal/datasources/mocks"
10+
"github.com/jbeshir/alignment-research-feed/internal/domain"
11+
"github.com/stretchr/testify/mock"
12+
"github.com/stretchr/testify/require"
13+
)
14+
15+
func testLogger() *slog.Logger {
16+
return slog.New(slog.DiscardHandler)
17+
}
18+
19+
func TestAddArticleToUserVector_Execute(t *testing.T) {
20+
testVector := []float32{0.1, 0.2, 0.3}
21+
22+
cases := []struct {
23+
name string
24+
userID string
25+
articleHashID string
26+
vector []float32
27+
vectorErr error
28+
addReturns bool
29+
addErr error
30+
wantAddCall bool
31+
}{
32+
{
33+
name: "adds_vector",
34+
userID: "user1",
35+
articleHashID: "article1",
36+
vector: testVector,
37+
addReturns: true,
38+
wantAddCall: true,
39+
},
40+
{
41+
name: "already_added_no_change",
42+
userID: "user1",
43+
articleHashID: "article1",
44+
vector: testVector,
45+
addReturns: false, // syncer returns false when already added
46+
wantAddCall: true,
47+
},
48+
{
49+
name: "no_vector_skips",
50+
userID: "user1",
51+
articleHashID: "article1",
52+
vector: nil,
53+
wantAddCall: false,
54+
},
55+
}
56+
57+
for _, tc := range cases {
58+
t.Run(tc.name, func(t *testing.T) {
59+
fetcher := mocks.NewMockArticleVectorFetcher(t)
60+
syncer := mocks.NewMockUserVectorSyncer(t)
61+
62+
fetcher.EXPECT().
63+
FetchArticleVector(mock.Anything, tc.articleHashID).
64+
Return(tc.vector, tc.vectorErr)
65+
66+
if tc.wantAddCall {
67+
syncer.EXPECT().
68+
AddArticleVectorToUser(mock.Anything, tc.userID, tc.articleHashID, tc.vector).
69+
Return(tc.addReturns, tc.addErr)
70+
}
71+
72+
cmd := &AddArticleToUserVector{
73+
ArticleVectorFetcher: fetcher,
74+
UserVectorSyncer: syncer,
75+
}
76+
77+
ctx := domain.ContextWithLogger(context.Background(), testLogger())
78+
err := cmd.Execute(ctx, tc.userID, tc.articleHashID)
79+
require.NoError(t, err)
80+
})
81+
}
82+
}
83+
84+
func TestAddArticleToUserVector_Execute_FetchError(t *testing.T) {
85+
fetcher := mocks.NewMockArticleVectorFetcher(t)
86+
syncer := mocks.NewMockUserVectorSyncer(t)
87+
88+
fetcher.EXPECT().
89+
FetchArticleVector(mock.Anything, "article1").
90+
Return(nil, errors.New("pinecone error"))
91+
92+
cmd := &AddArticleToUserVector{
93+
ArticleVectorFetcher: fetcher,
94+
UserVectorSyncer: syncer,
95+
}
96+
97+
ctx := domain.ContextWithLogger(context.Background(), testLogger())
98+
err := cmd.Execute(ctx, "user1", "article1")
99+
100+
// Fetch errors are logged but not returned
101+
require.NoError(t, err)
102+
}
103+
104+
func TestAddArticleToUserVector_Execute_AddError(t *testing.T) {
105+
fetcher := mocks.NewMockArticleVectorFetcher(t)
106+
syncer := mocks.NewMockUserVectorSyncer(t)
107+
108+
testVector := []float32{0.1}
109+
110+
fetcher.EXPECT().
111+
FetchArticleVector(mock.Anything, "article1").
112+
Return(testVector, nil)
113+
114+
syncer.EXPECT().
115+
AddArticleVectorToUser(mock.Anything, "user1", "article1", testVector).
116+
Return(false, errors.New("db error"))
117+
118+
cmd := &AddArticleToUserVector{
119+
ArticleVectorFetcher: fetcher,
120+
UserVectorSyncer: syncer,
121+
}
122+
123+
ctx := domain.ContextWithLogger(context.Background(), testLogger())
124+
err := cmd.Execute(ctx, "user1", "article1")
125+
126+
require.Error(t, err)
127+
require.Contains(t, err.Error(), "adding article vector to user")
128+
}
Lines changed: 160 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,160 @@
1+
package command
2+
3+
import (
4+
"context"
5+
"errors"
6+
"testing"
7+
8+
"github.com/jbeshir/alignment-research-feed/internal/datasources/mocks"
9+
"github.com/jbeshir/alignment-research-feed/internal/domain"
10+
"github.com/stretchr/testify/assert"
11+
"github.com/stretchr/testify/mock"
12+
"github.com/stretchr/testify/require"
13+
)
14+
15+
func TestDivideVector(t *testing.T) {
16+
cases := []struct {
17+
name string
18+
vector []float32
19+
divisor float32
20+
expected []float32
21+
}{
22+
{
23+
name: "simple_division",
24+
vector: []float32{2.0, 4.0, 6.0},
25+
divisor: 2.0,
26+
expected: []float32{1.0, 2.0, 3.0},
27+
},
28+
{
29+
name: "divide_by_one",
30+
vector: []float32{1.0, 2.0, 3.0},
31+
divisor: 1.0,
32+
expected: []float32{1.0, 2.0, 3.0},
33+
},
34+
{
35+
name: "empty_vector",
36+
vector: []float32{},
37+
divisor: 2.0,
38+
expected: []float32{},
39+
},
40+
{
41+
name: "fractional_result",
42+
vector: []float32{1.0, 2.0, 3.0},
43+
divisor: 2.0,
44+
expected: []float32{0.5, 1.0, 1.5},
45+
},
46+
}
47+
48+
for _, tc := range cases {
49+
t.Run(tc.name, func(t *testing.T) {
50+
result := divideVector(tc.vector, tc.divisor)
51+
assert.Equal(t, tc.expected, result)
52+
})
53+
}
54+
}
55+
56+
func TestRecommendArticles_Execute(t *testing.T) {
57+
cases := []struct {
58+
name string
59+
vectorSum []float32
60+
vectorCount int
61+
similar []domain.SimilarArticle
62+
articles []domain.Article
63+
expected []domain.Article
64+
wantErr bool
65+
errContains string
66+
}{
67+
{
68+
name: "no_vector_returns_nil",
69+
vectorSum: nil,
70+
vectorCount: 0,
71+
expected: nil,
72+
},
73+
{
74+
name: "zero_count_returns_nil",
75+
vectorSum: []float32{1.0, 2.0},
76+
vectorCount: 0,
77+
expected: nil,
78+
},
79+
{
80+
name: "successful_recommendation",
81+
vectorSum: []float32{2.0, 4.0, 6.0},
82+
vectorCount: 2,
83+
similar: []domain.SimilarArticle{
84+
{HashID: "rec1", Score: 0.9},
85+
{HashID: "rec2", Score: 0.8},
86+
},
87+
articles: []domain.Article{
88+
{HashID: "rec1", Title: "Recommended 1"},
89+
{HashID: "rec2", Title: "Recommended 2"},
90+
},
91+
expected: []domain.Article{
92+
{HashID: "rec1", Title: "Recommended 1"},
93+
{HashID: "rec2", Title: "Recommended 2"},
94+
},
95+
},
96+
{
97+
name: "no_similar_articles",
98+
vectorSum: []float32{1.0, 2.0},
99+
vectorCount: 1,
100+
similar: []domain.SimilarArticle{},
101+
expected: nil,
102+
},
103+
}
104+
105+
for _, tc := range cases {
106+
t.Run(tc.name, func(t *testing.T) {
107+
vectorSimilarity := mocks.NewMockSimilarArticlesByVectorLister(t)
108+
articleFetcher := mocks.NewMockArticleFetcher(t)
109+
userVectorGetter := mocks.NewMockUserVectorGetter(t)
110+
111+
userVectorGetter.EXPECT().
112+
GetUserVector(mock.Anything, "user1").
113+
Return(tc.vectorSum, tc.vectorCount, nil)
114+
115+
// Only expect similarity and fetch calls when we have a valid vector
116+
if tc.vectorSum != nil && tc.vectorCount > 0 {
117+
vectorSimilarity.EXPECT().
118+
ListSimilarArticlesByVector(mock.Anything, mock.Anything, mock.Anything, 10).
119+
Return(tc.similar, nil)
120+
121+
if len(tc.similar) > 0 {
122+
articleFetcher.EXPECT().
123+
FetchArticlesByID(mock.Anything, mock.Anything).
124+
Return(tc.articles, nil)
125+
}
126+
}
127+
128+
cmd := &RecommendArticles{
129+
VectorSimilarity: vectorSimilarity,
130+
ArticleFetcher: articleFetcher,
131+
UserVectorGetter: userVectorGetter,
132+
}
133+
134+
result, err := cmd.Execute(context.Background(), "user1", 10)
135+
136+
if tc.wantErr {
137+
require.Error(t, err)
138+
assert.Contains(t, err.Error(), tc.errContains)
139+
} else {
140+
require.NoError(t, err)
141+
assert.Equal(t, tc.expected, result)
142+
}
143+
})
144+
}
145+
}
146+
147+
func TestRecommendArticles_Execute_GetUserVectorError(t *testing.T) {
148+
userVectorGetter := mocks.NewMockUserVectorGetter(t)
149+
userVectorGetter.EXPECT().
150+
GetUserVector(mock.Anything, "user1").
151+
Return(nil, 0, errors.New("db error"))
152+
153+
cmd := &RecommendArticles{
154+
UserVectorGetter: userVectorGetter,
155+
}
156+
157+
_, err := cmd.Execute(context.Background(), "user1", 10)
158+
require.Error(t, err)
159+
assert.Contains(t, err.Error(), "getting user vector")
160+
}

0 commit comments

Comments
 (0)