Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions cmd/api/routes.go
Original file line number Diff line number Diff line change
Expand Up @@ -125,6 +125,7 @@ func (app *application) routes() http.Handler {
authGroup.DELETE("/tasks/:id", tasksHandler.DeleteTask)

/* Messages routes */
authGroup.GET("/messages/users", messagesHandler.GetMessageableUsers)
authGroup.GET("/messages/:userId", messagesHandler.GetMessagesBetweenUsers)

// WebSocket Connection
Expand Down
1 change: 1 addition & 0 deletions internal/adapters/postgresql/sqlc/querier.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

30 changes: 30 additions & 0 deletions internal/adapters/postgresql/sqlc/queries.sql
Original file line number Diff line number Diff line change
Expand Up @@ -192,6 +192,36 @@ WHERE (sender_id = $1 AND receiver_id = $2)
ORDER BY created_at DESC LIMIT $3
OFFSET $4;

-- name: GetMessageableUsers :many
WITH user_workspaces AS (
SELECT w.id
FROM workspaces w
WHERE w.user_id = sqlc.arg(logged_user_id)
OR EXISTS (SELECT 1
FROM workspace_members wm
WHERE wm.workspace_id = w.id
AND wm.user_id = sqlc.arg(logged_user_id))
)
SELECT count(*) OVER () AS total_count,
u.id,
u.first_name,
u.last_name,
u.email,
u.created_at,
u.updated_at
FROM users u
WHERE u.id <> sqlc.arg(logged_user_id)
AND (EXISTS (SELECT 1
FROM workspace_members owm
WHERE owm.user_id = u.id
AND owm.workspace_id IN (SELECT id FROM user_workspaces))
OR EXISTS (SELECT 1
FROM workspaces ow
WHERE ow.user_id = u.id
AND ow.id IN (SELECT id FROM user_workspaces)))
ORDER BY u.created_at DESC
LIMIT sqlc.arg(page_limit) OFFSET sqlc.arg(page_offset);

-- name: UpdateUserStripeCustomer :one
UPDATE users
SET stripe_customer_id = $2
Expand Down
75 changes: 75 additions & 0 deletions internal/adapters/postgresql/sqlc/queries.sql.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

22 changes: 15 additions & 7 deletions internal/auth/middleware.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,21 +4,29 @@ import (
"strings"

codeerror "gin-api-1/internal/codeerror"

"github.com/gin-gonic/gin"
)

func AuthenticationMiddleware(service Service) gin.HandlerFunc {
return func(c *gin.Context) {
authHeader := c.GetHeader("Authorization")
if authHeader == "" {
codeerror.HandleError(c, codeerror.New(codeerror.MissingToken, "Authorization header is missing"))
c.Abort()
return
token := ""
if authHeader != "" {
token = strings.TrimPrefix(authHeader, "Bearer ")
if token == authHeader {
codeerror.HandleError(c, codeerror.New(codeerror.InvalidToken, "Bearer token is invalid"))
c.Abort()
return
}
} else {
// Browsers cannot set custom headers on WebSocket connections, so
// WebSocket clients may authenticate via ?token=<JWT> instead.
token = c.Query("token")
}

token := strings.TrimPrefix(authHeader, "Bearer ")
if token == authHeader {
codeerror.HandleError(c, codeerror.New(codeerror.InvalidToken, "Bearer token is invalid"))
if token == "" {
codeerror.HandleError(c, codeerror.New(codeerror.MissingToken, "Authorization header is missing"))
c.Abort()
return
}
Expand Down
22 changes: 22 additions & 0 deletions internal/messages/handlers.go
Original file line number Diff line number Diff line change
Expand Up @@ -64,3 +64,25 @@ func (h *handler) GetMessagesBetweenUsers(c *gin.Context) {

c.JSON(http.StatusOK, response)
}

func (h *handler) GetMessageableUsers(c *gin.Context) {
loggedUser := c.MustGet("user").(auth.UserResponse)

page, err := strconv.Atoi(c.DefaultQuery("page", "1"))
if err != nil || page < 1 {
page = 1
}

pageSize, err := strconv.Atoi(c.DefaultQuery("pageSize", strconv.Itoa(DefaultPageSize)))
if err != nil || pageSize < 1 {
pageSize = DefaultPageSize
}

response, err := h.service.GetMessageableUsers(c, loggedUser.ID, page, pageSize)
if err != nil {
codeerror.HandleError(c, err)
return
}

c.JSON(http.StatusOK, response)
}
127 changes: 127 additions & 0 deletions internal/messages/service_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ type mockMessagesRepository struct {
getUserByIdFunc func(ctx context.Context, id pgtype.UUID) (repo.User, error)
createMessageFunc func(ctx context.Context, arg repo.CreateMessageParams) (repo.Message, error)
getMessagesBetweenUsersFunc func(ctx context.Context, arg repo.GetMessagesBetweenUsersParams) ([]repo.GetMessagesBetweenUsersRow, error)
getMessageableUsersFunc func(ctx context.Context, arg repo.GetMessageableUsersParams) ([]repo.GetMessageableUsersRow, error)
}

func (m *mockMessagesRepository) GetUserById(ctx context.Context, id pgtype.UUID) (repo.User, error) {
Expand All @@ -29,6 +30,10 @@ func (m *mockMessagesRepository) GetMessagesBetweenUsers(ctx context.Context, ar
return m.getMessagesBetweenUsersFunc(ctx, arg)
}

func (m *mockMessagesRepository) GetMessageableUsers(ctx context.Context, arg repo.GetMessageableUsersParams) ([]repo.GetMessageableUsersRow, error) {
return m.getMessageableUsersFunc(ctx, arg)
}

func newTestMessage(id, senderID, receiverID uuid.UUID, content string) repo.Message {
return repo.Message{
ID: pgtype.UUID{Bytes: id, Valid: true},
Expand Down Expand Up @@ -386,3 +391,125 @@ func TestGetMessagesBetweenUsers(t *testing.T) {
})
}
}

func TestGetMessageableUsers(t *testing.T) {
ctx := context.Background()
loggedUserID := uuid.New()
otherUserID := uuid.New()

otherUser := repo.GetMessageableUsersRow{
TotalCount: 1,
ID: pgtype.UUID{Bytes: otherUserID, Valid: true},
FirstName: "Jane",
LastName: "Doe",
Email: "jane@example.com",
CreatedAt: pgtype.Timestamptz{Time: time.Now(), Valid: true},
UpdatedAt: pgtype.Timestamptz{Time: time.Now(), Valid: true},
}

tests := []struct {
name string
userID string
repo *mockMessagesRepository
wantErr bool
wantLen int
wantPages int
}{
{
name: "success",
userID: loggedUserID.String(),
repo: &mockMessagesRepository{
getMessageableUsersFunc: func(ctx context.Context, arg repo.GetMessageableUsersParams) ([]repo.GetMessageableUsersRow, error) {
if arg.LoggedUserID.Bytes != loggedUserID {
t.Errorf("logged user ID = %v, want %v", arg.LoggedUserID.Bytes, loggedUserID)
}
return []repo.GetMessageableUsersRow{otherUser}, nil
},
},
wantErr: false,
wantLen: 1,
wantPages: 1,
},
{
name: "invalid uuid",
userID: "not-a-uuid",
repo: &mockMessagesRepository{
getMessageableUsersFunc: func(ctx context.Context, arg repo.GetMessageableUsersParams) ([]repo.GetMessageableUsersRow, error) {
t.Fatal("repo should not be called for an invalid UUID")
return nil, nil
},
},
wantErr: true,
},
{
name: "empty result",
userID: loggedUserID.String(),
repo: &mockMessagesRepository{
getMessageableUsersFunc: func(ctx context.Context, arg repo.GetMessageableUsersParams) ([]repo.GetMessageableUsersRow, error) {
return []repo.GetMessageableUsersRow{}, nil
},
},
wantErr: false,
wantLen: 0,
wantPages: 0,
},
{
name: "calculates total pages",
userID: loggedUserID.String(),
repo: &mockMessagesRepository{
getMessageableUsersFunc: func(ctx context.Context, arg repo.GetMessageableUsersParams) ([]repo.GetMessageableUsersRow, error) {
return []repo.GetMessageableUsersRow{otherUser}, nil
},
},
wantErr: false,
wantLen: 1,
wantPages: 1,
},
{
name: "repo failure",
userID: loggedUserID.String(),
repo: &mockMessagesRepository{
getMessageableUsersFunc: func(ctx context.Context, arg repo.GetMessageableUsersParams) ([]repo.GetMessageableUsersRow, error) {
return nil, pgx.ErrTxClosed
},
},
wantErr: true,
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
service := &svc{repo: tt.repo}
result, err := service.GetMessageableUsers(ctx, tt.userID, 1, 10)

if tt.wantErr {
if err == nil {
t.Fatal("expected error, got nil")
}
return
}

if err != nil {
t.Fatalf("unexpected error: %v", err)
}

if len(result.Users) != tt.wantLen {
t.Errorf("users count = %d, want %d", len(result.Users), tt.wantLen)
}

if tt.wantLen > 0 {
if result.Users[0].ID != otherUserID.String() {
t.Errorf("user ID = %s, want %s", result.Users[0].ID, otherUserID)
}
}

if result.Pagination.Page != 1 || result.Pagination.PageSize != 10 {
t.Errorf("pagination = %+v, want page=1 pageSize=10", result.Pagination)
}

if result.Pagination.TotalPages != tt.wantPages {
t.Errorf("total pages = %d, want %d", result.Pagination.TotalPages, tt.wantPages)
}
})
}
}
Loading
Loading