Skip to content

Commit c7482d9

Browse files
committed
fix(api): persist new unread threads
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: bf00a0ac-e11f-4015-b295-3cdd9b491229
1 parent 3fc71ba commit c7482d9

2 files changed

Lines changed: 112 additions & 1 deletion

File tree

api/pkg/repositories/gorm_message_thread_repository.go

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -111,7 +111,24 @@ func (repository *gormMessageThreadRepository) Store(ctx context.Context, thread
111111
ctx, span := repository.tracer.Start(ctx)
112112
defer span.End()
113113

114-
if err := repository.db.WithContext(ctx).Clauses(clause.OnConflict{DoNothing: true}).Create(thread).Error; err != nil {
114+
isRead := thread.IsRead
115+
err := repository.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
116+
result := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(thread)
117+
thread.IsRead = isRead
118+
if result.Error != nil {
119+
return result.Error
120+
}
121+
if result.RowsAffected == 0 || isRead {
122+
return nil
123+
}
124+
125+
return tx.Model(&entities.MessageThread{}).
126+
Where("user_id = ?", thread.UserID).
127+
Where("id = ?", thread.ID).
128+
UpdateColumn("is_read", false).
129+
Error
130+
})
131+
if err != nil {
115132
msg := fmt.Sprintf("cannot save message thread with ID [%s]", thread.ID)
116133
return repository.tracer.WrapErrorSpan(span, stacktrace.Propagate(err, msg))
117134
}

api/pkg/repositories/gorm_message_thread_repository_test.go

Lines changed: 94 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,108 @@
11
package repositories
22

33
import (
4+
"context"
5+
"database/sql"
6+
"database/sql/driver"
7+
"errors"
8+
"strings"
49
"testing"
510
"time"
611

712
"github.com/NdoleStudio/httpsms/pkg/entities"
13+
"github.com/NdoleStudio/httpsms/pkg/telemetry"
814
"github.com/google/uuid"
915
"github.com/stretchr/testify/assert"
16+
"github.com/stretchr/testify/require"
17+
"go.opentelemetry.io/otel/trace"
18+
"gorm.io/driver/postgres"
19+
"gorm.io/gorm"
1020
)
1121

22+
type messageThreadTestStatement struct {
23+
query string
24+
args []any
25+
}
26+
27+
type messageThreadTestConnPool struct {
28+
statements []messageThreadTestStatement
29+
}
30+
31+
func (messageThreadTestConnPool) PrepareContext(context.Context, string) (*sql.Stmt, error) {
32+
return nil, errors.New("unexpected PrepareContext")
33+
}
34+
35+
func (pool *messageThreadTestConnPool) ExecContext(_ context.Context, query string, args ...any) (sql.Result, error) {
36+
pool.statements = append(pool.statements, messageThreadTestStatement{
37+
query: query,
38+
args: append([]any(nil), args...),
39+
})
40+
return driver.RowsAffected(1), nil
41+
}
42+
43+
func (messageThreadTestConnPool) QueryContext(context.Context, string, ...any) (*sql.Rows, error) {
44+
return nil, errors.New("unexpected QueryContext")
45+
}
46+
47+
func (messageThreadTestConnPool) QueryRowContext(context.Context, string, ...any) *sql.Row {
48+
return &sql.Row{}
49+
}
50+
51+
func (pool *messageThreadTestConnPool) BeginTx(context.Context, *sql.TxOptions) (gorm.ConnPool, error) {
52+
return pool, nil
53+
}
54+
55+
func (*messageThreadTestConnPool) Commit() error {
56+
return nil
57+
}
58+
59+
func (*messageThreadTestConnPool) Rollback() error {
60+
return nil
61+
}
62+
63+
type messageThreadTestLogger struct{}
64+
65+
func (logger *messageThreadTestLogger) Error(error) {}
66+
func (logger *messageThreadTestLogger) WithService(string) telemetry.Logger { return logger }
67+
68+
func (logger *messageThreadTestLogger) WithString(string, string) telemetry.Logger { return logger }
69+
70+
func (logger *messageThreadTestLogger) WithSpan(trace.SpanContext) telemetry.Logger { return logger }
71+
func (logger *messageThreadTestLogger) Trace(string) {}
72+
func (logger *messageThreadTestLogger) Info(string) {}
73+
func (logger *messageThreadTestLogger) Warn(error) {}
74+
func (logger *messageThreadTestLogger) Debug(string) {}
75+
func (logger *messageThreadTestLogger) Fatal(error) {}
76+
func (logger *messageThreadTestLogger) Printf(string, ...interface{}) {}
77+
78+
func TestMessageThreadStorePreservesExplicitUnreadState(t *testing.T) {
79+
pool := &messageThreadTestConnPool{}
80+
db, err := gorm.Open(
81+
postgres.New(postgres.Config{
82+
Conn: pool,
83+
WithoutReturning: true,
84+
}),
85+
&gorm.Config{DisableAutomaticPing: true},
86+
)
87+
require.NoError(t, err)
88+
89+
logger := &messageThreadTestLogger{}
90+
repository := NewGormMessageThreadRepository(logger, telemetry.NewOtelLogger("test", logger), db)
91+
thread := &entities.MessageThread{
92+
ID: uuid.New(),
93+
IsRead: false,
94+
}
95+
96+
require.NoError(t, repository.Store(context.Background(), thread))
97+
assert.False(t, thread.IsRead)
98+
99+
require.NotEmpty(t, pool.statements)
100+
update := pool.statements[len(pool.statements)-1]
101+
assert.True(t, strings.HasPrefix(update.query, `UPDATE "message_threads"`))
102+
assert.Contains(t, update.query, `"is_read"=$1`)
103+
assert.Contains(t, update.args, false)
104+
}
105+
12106
func TestMessageThreadActivityUpdatesOwnOnlyMessageColumns(t *testing.T) {
13107
messageID := uuid.New()
14108
updates := messageThreadActivityUpdates(MessageThreadActivityUpdate{

0 commit comments

Comments
 (0)