|
1 | 1 | package repositories |
2 | 2 |
|
3 | 3 | import ( |
| 4 | + "context" |
| 5 | + "database/sql" |
| 6 | + "database/sql/driver" |
| 7 | + "errors" |
| 8 | + "strings" |
4 | 9 | "testing" |
5 | 10 | "time" |
6 | 11 |
|
7 | 12 | "github.com/NdoleStudio/httpsms/pkg/entities" |
| 13 | + "github.com/NdoleStudio/httpsms/pkg/telemetry" |
8 | 14 | "github.com/google/uuid" |
9 | 15 | "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" |
10 | 20 | ) |
11 | 21 |
|
| 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 | + |
12 | 106 | func TestMessageThreadActivityUpdatesOwnOnlyMessageColumns(t *testing.T) { |
13 | 107 | messageID := uuid.New() |
14 | 108 | updates := messageThreadActivityUpdates(MessageThreadActivityUpdate{ |
|
0 commit comments