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
24 changes: 20 additions & 4 deletions server/config/pgx.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"context"
"fmt"
"log"
"net"
"strings"
"time"

Expand All @@ -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"
}

Expand Down
74 changes: 73 additions & 1 deletion server/controllers/messages.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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})
}
5 changes: 5 additions & 0 deletions server/migrations/012_add_edited_at_to_messages_up.sql
Original file line number Diff line number Diff line change
@@ -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;
2 changes: 1 addition & 1 deletion server/models/message.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"`
Expand All @@ -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"`
}
4 changes: 2 additions & 2 deletions server/repository/groups.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
}
Expand Down
61 changes: 57 additions & 4 deletions server/repository/message.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{}
Expand Down Expand Up @@ -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
Expand All @@ -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
}
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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
}
10 changes: 6 additions & 4 deletions server/routes/messages.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Loading