diff --git a/cmd/api/routes.go b/cmd/api/routes.go index 0362188..8f7cb3c 100644 --- a/cmd/api/routes.go +++ b/cmd/api/routes.go @@ -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 diff --git a/internal/adapters/postgresql/sqlc/querier.go b/internal/adapters/postgresql/sqlc/querier.go index f7b3ce9..a0c13ee 100644 --- a/internal/adapters/postgresql/sqlc/querier.go +++ b/internal/adapters/postgresql/sqlc/querier.go @@ -28,6 +28,7 @@ type Querier interface { DeleteWorkspace(ctx context.Context, arg DeleteWorkspaceParams) error GetAccessibleWorkspaceByID(ctx context.Context, arg GetAccessibleWorkspaceByIDParams) (Workspace, error) GetMemberFromWorkspace(ctx context.Context, arg GetMemberFromWorkspaceParams) (WorkspaceMember, error) + GetMessageableUsers(ctx context.Context, arg GetMessageableUsersParams) ([]GetMessageableUsersRow, error) GetMessagesBetweenUsers(ctx context.Context, arg GetMessagesBetweenUsersParams) ([]GetMessagesBetweenUsersRow, error) GetProjectById(ctx context.Context, id pgtype.UUID) (Project, error) GetProjectIntegration(ctx context.Context, projectID pgtype.UUID) (ProjectIntegration, error) diff --git a/internal/adapters/postgresql/sqlc/queries.sql b/internal/adapters/postgresql/sqlc/queries.sql index fddb84e..e06361e 100644 --- a/internal/adapters/postgresql/sqlc/queries.sql +++ b/internal/adapters/postgresql/sqlc/queries.sql @@ -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 diff --git a/internal/adapters/postgresql/sqlc/queries.sql.go b/internal/adapters/postgresql/sqlc/queries.sql.go index 4715508..e2a0276 100644 --- a/internal/adapters/postgresql/sqlc/queries.sql.go +++ b/internal/adapters/postgresql/sqlc/queries.sql.go @@ -506,6 +506,81 @@ func (q *Queries) GetMemberFromWorkspace(ctx context.Context, arg GetMemberFromW return i, err } +const getMessageableUsers = `-- name: GetMessageableUsers :many +WITH user_workspaces AS ( + SELECT w.id + FROM workspaces w + WHERE w.user_id = $1 + OR EXISTS (SELECT 1 + FROM workspace_members wm + WHERE wm.workspace_id = w.id + AND wm.user_id = $1) +) +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 <> $1 + 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 $3 OFFSET $2 +` + +type GetMessageableUsersParams struct { + LoggedUserID pgtype.UUID `json:"logged_user_id"` + PageOffset int32 `json:"page_offset"` + PageLimit int32 `json:"page_limit"` +} + +type GetMessageableUsersRow struct { + TotalCount int64 `json:"total_count"` + ID pgtype.UUID `json:"id"` + FirstName string `json:"first_name"` + LastName string `json:"last_name"` + Email string `json:"email"` + CreatedAt pgtype.Timestamptz `json:"created_at"` + UpdatedAt pgtype.Timestamptz `json:"updated_at"` +} + +func (q *Queries) GetMessageableUsers(ctx context.Context, arg GetMessageableUsersParams) ([]GetMessageableUsersRow, error) { + rows, err := q.db.Query(ctx, getMessageableUsers, arg.LoggedUserID, arg.PageOffset, arg.PageLimit) + if err != nil { + return nil, err + } + defer rows.Close() + var items []GetMessageableUsersRow + for rows.Next() { + var i GetMessageableUsersRow + if err := rows.Scan( + &i.TotalCount, + &i.ID, + &i.FirstName, + &i.LastName, + &i.Email, + &i.CreatedAt, + &i.UpdatedAt, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + const getMessagesBetweenUsers = `-- name: GetMessagesBetweenUsers :many SELECT count(*) OVER () AS total_count, id, sender_id, diff --git a/internal/auth/middleware.go b/internal/auth/middleware.go index 7352ee2..5d1bd4a 100644 --- a/internal/auth/middleware.go +++ b/internal/auth/middleware.go @@ -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= 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 } diff --git a/internal/messages/handlers.go b/internal/messages/handlers.go index dfd73b5..15e180e 100644 --- a/internal/messages/handlers.go +++ b/internal/messages/handlers.go @@ -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) +} diff --git a/internal/messages/service_test.go b/internal/messages/service_test.go index 04be7da..88fa13f 100644 --- a/internal/messages/service_test.go +++ b/internal/messages/service_test.go @@ -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) { @@ -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}, @@ -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) + } + }) + } +} diff --git a/internal/messages/services.go b/internal/messages/services.go index be741c7..f52845b 100644 --- a/internal/messages/services.go +++ b/internal/messages/services.go @@ -3,6 +3,7 @@ package messages import ( "context" repo "gin-api-1/internal/adapters/postgresql/sqlc" + "gin-api-1/internal/auth" "gin-api-1/internal/codeerror" "github.com/google/uuid" @@ -13,12 +14,14 @@ import ( type Service interface { CreateMessage(ctx context.Context, senderID string, payload CreateMessagePayload) (repo.Message, error) GetMessagesBetweenUsers(ctx context.Context, loggedUserID, otherUserID string, page, pageSize int) (getMessagesResponse, error) + GetMessageableUsers(ctx context.Context, loggedUserID string, page, pageSize int) (getMessageableUsersResponse, error) } type messagesRepository interface { GetUserById(ctx context.Context, id pgtype.UUID) (repo.User, error) CreateMessage(ctx context.Context, arg repo.CreateMessageParams) (repo.Message, error) GetMessagesBetweenUsers(ctx context.Context, arg repo.GetMessagesBetweenUsersParams) ([]repo.GetMessagesBetweenUsersRow, error) + GetMessageableUsers(ctx context.Context, arg repo.GetMessageableUsersParams) ([]repo.GetMessageableUsersRow, error) } type svc struct { @@ -121,3 +124,47 @@ func (s *svc) GetMessagesBetweenUsers(ctx context.Context, loggedUserID, otherUs }, }, nil } + +func (s *svc) GetMessageableUsers(ctx context.Context, loggedUserID string, page, pageSize int) (getMessageableUsersResponse, error) { + loggedUUID, err := uuid.Parse(loggedUserID) + if err != nil { + return getMessageableUsersResponse{}, codeerror.New(codeerror.InvalidUUID, "User ID is not a valid UUID") + } + + rows, err := s.repo.GetMessageableUsers(ctx, repo.GetMessageableUsersParams{ + LoggedUserID: pgtype.UUID{Bytes: loggedUUID, Valid: true}, + PageOffset: int32((page - 1) * pageSize), + PageLimit: int32(pageSize), + }) + if err != nil { + return getMessageableUsersResponse{}, codeerror.Wrap(codeerror.StatusInternalServerError, "Failed to fetch messageable users", err) + } + + users := make([]auth.UserResponse, 0, len(rows)) + + var total int64 + if len(rows) > 0 { + total = rows[0].TotalCount + } + + for _, row := range rows { + users = append(users, auth.UserResponse{ + ID: row.ID.String(), + FirstName: row.FirstName, + LastName: row.LastName, + Email: row.Email, + CreatedAt: row.CreatedAt, + UpdatedAt: row.UpdatedAt, + }) + } + + return getMessageableUsersResponse{ + Users: users, + Pagination: paginationResponse{ + Page: page, + PageSize: pageSize, + Total: total, + TotalPages: (int(total) + pageSize - 1) / pageSize, + }, + }, nil +} diff --git a/internal/messages/types.go b/internal/messages/types.go index 6941a8d..020de82 100644 --- a/internal/messages/types.go +++ b/internal/messages/types.go @@ -1,6 +1,9 @@ package messages -import "time" +import ( + "gin-api-1/internal/auth" + "time" +) const DefaultPageSize = 10 @@ -28,3 +31,8 @@ type getMessagesResponse struct { Messages []MessageResponse `json:"messages"` Pagination paginationResponse `json:"pagination"` } + +type getMessageableUsersResponse struct { + Users []auth.UserResponse `json:"users"` + Pagination paginationResponse `json:"pagination"` +}