From f6d1ded4c03bcc675e11c2663d935c0075c11714 Mon Sep 17 00:00:00 2001 From: Suman Biswas Date: Mon, 10 Aug 2026 18:50:40 +0530 Subject: [PATCH 1/2] pgx fix --- server/config/pgx.go | 24 +++++++++++++++---- .../012_add_edited_at_to_messages_up.sql | 5 ++++ 2 files changed, 25 insertions(+), 4 deletions(-) create mode 100644 server/migrations/012_add_edited_at_to_messages_up.sql diff --git a/server/config/pgx.go b/server/config/pgx.go index 3ec991e..dd515ee 100644 --- a/server/config/pgx.go +++ b/server/config/pgx.go @@ -4,6 +4,7 @@ import ( "context" "fmt" "log" + "net" "strings" "time" @@ -13,14 +14,29 @@ import ( var DB *pgxpool.Pool -func isLocalHost(host string) bool { - h := strings.ToLower(host) - return h == "localhost" || h == "127.0.0.1" || h == "::1" || strings.HasPrefix(h, "172.") +func isPrivateOrLocal(host string) bool { + h := strings.ToLower(strings.TrimSpace(host)) + + // Local hostnames and Docker container service names + if h == "localhost" || h == "postgres" || h == "db" || h == "host.docker.internal" { + return true + } + + // Parse IP and check RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16) and Loopbacks + ip := net.ParseIP(h) + if ip != nil { + return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsUnspecified() + } + + // Fallback string prefix checks + return strings.HasPrefix(h, "10.") || + strings.HasPrefix(h, "192.168.") || + strings.HasPrefix(h, "127.") } func InitDatabase() { sslMode := "require" - if isLocalHost(env.POSTGRES_HOST) { + if isPrivateOrLocal(env.POSTGRES_HOST) { sslMode = "disable" } diff --git a/server/migrations/012_add_edited_at_to_messages_up.sql b/server/migrations/012_add_edited_at_to_messages_up.sql new file mode 100644 index 0000000..159a558 --- /dev/null +++ b/server/migrations/012_add_edited_at_to_messages_up.sql @@ -0,0 +1,5 @@ +ALTER TABLE messages ADD COLUMN IF NOT EXISTS edited_at TIMESTAMPTZ; + +CREATE INDEX IF NOT EXISTS idx_messages_edited_at +ON messages(edited_at) +WHERE edited_at IS NOT NULL; \ No newline at end of file From edffea98055d533e99933d7e457b16d14185e725 Mon Sep 17 00:00:00 2001 From: Suman Biswas Date: Mon, 10 Aug 2026 19:51:36 +0530 Subject: [PATCH 2/2] edit message --- server/controllers/messages.go | 74 +++++++++++++++++++++++++++++++++- server/models/message.go | 2 +- server/repository/groups.go | 4 +- server/repository/message.go | 61 ++++++++++++++++++++++++++-- server/routes/messages.go | 10 +++-- 5 files changed, 139 insertions(+), 12 deletions(-) diff --git a/server/controllers/messages.go b/server/controllers/messages.go index ea0eb74..ab9de5d 100644 --- a/server/controllers/messages.go +++ b/server/controllers/messages.go @@ -4,12 +4,14 @@ import ( "encoding/json" "net/http" "strconv" + "strings" "time" "github.com/commandlinecoding/elephant/server/middlewares" "github.com/commandlinecoding/elephant/server/models" "github.com/commandlinecoding/elephant/server/repository" "github.com/commandlinecoding/elephant/server/services" + "github.com/go-chi/chi/v5" ) type SendMessageReq struct { @@ -72,7 +74,7 @@ func HandleGetChatHistory(w http.ResponseWriter, r *http.Request) { limit = 50 } - before := time.Now() + before := time.Now().Add(time.Minute) if beforeStr != "" { if t, err := time.Parse(time.RFC3339, beforeStr); err == nil { before = t @@ -142,3 +144,73 @@ func HandleMarkRead(w http.ResponseWriter, r *http.Request) { _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: true, Data: "Target messages marked read successfully"}) } + +type EditMessageReq struct { + Content string `json:"content"` +} + +func HandleEditMessage(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + + uid, ok := r.Context().Value(middlewares.UserIDKey).(string) + if !ok { + w.WriteHeader(http.StatusUnauthorized) + _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: false, Error: "Unauthorized context"}) + return + } + + msgID := chi.URLParam(r, "id") + if msgID == "" { + w.WriteHeader(http.StatusBadRequest) + _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: false, Error: "Message ID required"}) + return + } + + var req EditMessageReq + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + w.WriteHeader(http.StatusBadRequest) + _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: false, Error: "Invalid json payload structure"}) + return + } + + content := strings.TrimSpace(req.Content) + if content == "" { + w.WriteHeader(http.StatusBadRequest) + _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: false, Error: "Content cannot be empty"}) + return + } + + repo := repository.NewMessageRepository() + msg, err := repo.GetMessageByID(r.Context(), msgID) + if err != nil { + if strings.Contains(err.Error(), "db error") { + w.WriteHeader(http.StatusInternalServerError) + _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: false, Error: err.Error()}) + return + } + w.WriteHeader(http.StatusNotFound) + _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: false, Error: "Message not found"}) + return + } + + if msg.SenderID != uid { + w.WriteHeader(http.StatusForbidden) + _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: false, Error: "You can only edit your own messages"}) + return + } + + if time.Since(msg.CreatedAt) > 15*time.Minute { + w.WriteHeader(http.StatusForbidden) + _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: false, Error: "Message can no longer be edited (15 minute window expired)"}) + return + } + + updatedMsg, err := repo.UpdateMessage(r.Context(), msgID, content) + if err != nil { + w.WriteHeader(http.StatusInternalServerError) + _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: false, Error: "Failed to update message"}) + return + } + + _ = json.NewEncoder(w).Encode(models.JSONResponse{Success: true, Data: updatedMsg}) +} diff --git a/server/models/message.go b/server/models/message.go index 9bf3583..8d77ce6 100644 --- a/server/models/message.go +++ b/server/models/message.go @@ -2,7 +2,6 @@ package models import "time" -// QuotedMessage represents a lightweight snapshot of the parent message being replied to type QuotedMessage struct { ID string `json:"id"` SenderID string `json:"sender_id"` @@ -17,6 +16,7 @@ type Message struct { Content string `json:"content"` CreatedAt time.Time `json:"created_at"` IsRead bool `json:"is_read"` + EditedAt *time.Time `json:"edited_at,omitempty"` ReplyToMessageID *string `json:"reply_to_message_id,omitempty"` QuotedMessage *QuotedMessage `json:"quoted_message,omitempty"` } diff --git a/server/repository/groups.go b/server/repository/groups.go index c76cf91..6ce4001 100644 --- a/server/repository/groups.go +++ b/server/repository/groups.go @@ -106,7 +106,7 @@ func (r *GroupRepository) GetGroupMembers(ctx context.Context, groupID string) ( func (r *GroupRepository) GetGroupMessages(ctx context.Context, groupID string, before time.Time, limit int) ([]models.Message, error) { query := ` SELECT - m.id, m.sender_id, m.content, m.created_at, m.reply_to_message_id, + m.id, m.sender_id, m.content, m.created_at, m.edited_at, m.reply_to_message_id, q.sender_id AS quoted_sender_id, q.content AS quoted_content FROM messages m LEFT JOIN messages q ON m.reply_to_message_id = q.id @@ -125,7 +125,7 @@ func (r *GroupRepository) GetGroupMessages(ctx context.Context, groupID string, var m models.Message var qSender, qContent *string - err := rows.Scan(&m.ID, &m.SenderID, &m.Content, &m.CreatedAt, &m.ReplyToMessageID, &qSender, &qContent) + err := rows.Scan(&m.ID, &m.SenderID, &m.Content, &m.CreatedAt, &m.EditedAt, &m.ReplyToMessageID, &qSender, &qContent) if err != nil { return nil, err } diff --git a/server/repository/message.go b/server/repository/message.go index 2a7d4ef..98ff916 100644 --- a/server/repository/message.go +++ b/server/repository/message.go @@ -2,10 +2,13 @@ package repository import ( "context" + "errors" + "fmt" "time" "github.com/commandlinecoding/elephant/server/config" "github.com/commandlinecoding/elephant/server/models" + "github.com/jackc/pgx/v5" ) type MessageRepository struct{} @@ -38,7 +41,7 @@ func (r *MessageRepository) CreateMessage(ctx context.Context, senderID, receive func (r *MessageRepository) GetChatHistory(ctx context.Context, userA, userB string, before time.Time, limit int) ([]models.Message, error) { query := ` SELECT - m.id, m.sender_id, m.receiver_id, m.content, m.created_at, m.is_read, m.reply_to_message_id, + m.id, m.sender_id, m.receiver_id, m.content, m.created_at, m.is_read, m.edited_at, m.reply_to_message_id, q.sender_id AS quoted_sender_id, q.content AS quoted_content FROM messages m LEFT JOIN messages q ON m.reply_to_message_id = q.id @@ -58,7 +61,7 @@ func (r *MessageRepository) GetChatHistory(ctx context.Context, userA, userB str var m models.Message var qSender, qContent *string - err := rows.Scan(&m.ID, &m.SenderID, &m.ReceiverID, &m.Content, &m.CreatedAt, &m.IsRead, &m.ReplyToMessageID, &qSender, &qContent) + err := rows.Scan(&m.ID, &m.SenderID, &m.ReceiverID, &m.Content, &m.CreatedAt, &m.IsRead, &m.EditedAt, &m.ReplyToMessageID, &qSender, &qContent) if err != nil { return nil, err } @@ -88,7 +91,6 @@ func (r *MessageRepository) GetMessagePreview(ctx context.Context, messageID str func (r *MessageRepository) GetConversations(ctx context.Context, userID string) ([]models.Conversation, error) { query := ` WITH raw_conversations AS ( - -- Enclosing Part A inside explicit parentheses to isolate its internal ORDER BY (SELECT DISTINCT ON (CASE WHEN sender_id = $1::uuid THEN receiver_id ELSE sender_id END) CASE WHEN sender_id = $1::uuid THEN receiver_id ELSE sender_id END AS chat_user_id, 'direct' AS type, @@ -103,7 +105,6 @@ func (r *MessageRepository) GetConversations(ctx context.Context, userID string) UNION ALL - -- Enclosing Part B inside explicit parentheses to isolate its internal ORDER BY (SELECT DISTINCT ON (m.group_id) NULL::UUID AS chat_user_id, 'group' AS type, @@ -205,3 +206,55 @@ func (r *MessageRepository) UpdateGroupLastRead(ctx context.Context, groupID, us _, err := config.DB.Exec(ctx, query, groupID, userID) return err } + +func (r *MessageRepository) GetMessageByID(ctx context.Context, id string) (*models.Message, error) { + query := ` + SELECT id, sender_id, COALESCE(receiver_id::text, '') as receiver_id, COALESCE(group_id::text, '') as group_id, + content, created_at, edited_at, reply_to_message_id + FROM messages + WHERE id = $1::uuid; + ` + var msg models.Message + err := config.DB.QueryRow(ctx, query, id).Scan( + &msg.ID, + &msg.SenderID, + &msg.ReceiverID, + &msg.GroupID, + &msg.Content, + &msg.CreatedAt, + &msg.EditedAt, + &msg.ReplyToMessageID, + ) + if err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return nil, err + } + return nil, fmt.Errorf("db error: %w", err) + } + return &msg, nil +} + +func (r *MessageRepository) UpdateMessage(ctx context.Context, id, content string) (*models.Message, error) { + query := ` + UPDATE messages + SET content = $1, edited_at = NOW() + WHERE id = $2::uuid + RETURNING id, sender_id, COALESCE(receiver_id::text, '') as receiver_id, COALESCE(group_id::text, '') as group_id, + content, created_at, edited_at, reply_to_message_id; + ` + var msg models.Message + err := config.DB.QueryRow(ctx, query, content, id).Scan( + &msg.ID, + &msg.SenderID, + &msg.ReceiverID, + &msg.GroupID, + &msg.Content, + &msg.CreatedAt, + &msg.EditedAt, + &msg.ReplyToMessageID, + ) + if err != nil { + return nil, err + } + return &msg, nil +} diff --git a/server/routes/messages.go b/server/routes/messages.go index f01f5eb..4ca4f7c 100644 --- a/server/routes/messages.go +++ b/server/routes/messages.go @@ -8,8 +8,10 @@ import ( func MessageRoute(router chi.Router) { router.Use(middlewares.AuthGuard) - router.Post("/", controllers.HandleSendMessage) // Resolves to POST /api/messages - router.Get("/", controllers.HandleGetChatHistory) // Resolves to GET /api/messages?with=UUID - router.Get("/conversations", controllers.HandleGetConversations) // GET /api/messages/conversations - router.Post("/read", controllers.HandleMarkRead) // POST /api/messages/read + + router.Post("/", controllers.HandleSendMessage) // POST /api/messages + router.Get("/", controllers.HandleGetChatHistory) // GET /api/messages?with=UUID + router.Get("/conversations", controllers.HandleGetConversations) // GET /api/messages/conversations + router.Put("/{id}", controllers.HandleEditMessage) // PUT /api/messages/:id + router.Post("/read", controllers.HandleMarkRead) // POST /api/messages/read }