From e415a13211316cb804752c18833a10057bc0994c Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Wed, 27 Nov 2024 11:14:36 +0200 Subject: Reimplemented existing server APIs --- internal/server/api/api.go | 42 ++++++++++++++++++------------------------ internal/server/server.go | 28 +++++++++++----------------- 2 files changed, 29 insertions(+), 41 deletions(-) (limited to 'internal') diff --git a/internal/server/api/api.go b/internal/server/api/api.go index f60e0ec..6c8d51a 100644 --- a/internal/server/api/api.go +++ b/internal/server/api/api.go @@ -11,14 +11,10 @@ import ( "github.com/kyren223/eko/internal/data" "github.com/kyren223/eko/internal/packet" "github.com/kyren223/eko/internal/server/session" - "github.com/kyren223/eko/pkg/assert" "github.com/kyren223/eko/pkg/snowflake" ) -func SendMessage(ctx context.Context, request *packet.SendMessage) packet.Payload { - sess, ok := session.FromContext(ctx) - assert.Assert(ok, "context in process packet should always have a session") - +func SendMessage(ctx context.Context, sess *session.Session, request *packet.SendMessage) packet.Payload { if (request.ReceiverID != nil) == (request.FrequencyID != nil) { return &packet.Error{Error: "either receiver id or frequency id must exist"} } @@ -28,27 +24,38 @@ func SendMessage(ctx context.Context, request *packet.SendMessage) packet.Payloa return &packet.Error{Error: "message content must not be blank"} } - node := sess.Manager().Node() - queries := data.New(db) message, err := queries.CreateMessage(ctx, data.CreateMessageParams{ - ID: node.Generate(), + ID: sess.Manager().Node().Generate(), SenderID: sess.ID(), Content: content, FrequencyID: request.FrequencyID, ReceiverID: request.ReceiverID, }) if err != nil { - log.Println(sess.Addr(), "SendMessage database error:", err) + log.Println(sess.Addr(), "database error:", err, "in SendMessage") return &packet.Error{Error: "internal server error"} } return &packet.MessagesInfo{Messages: []data.Message{message}} } -func GetMessages(ctx context.Context, request *packet.RequestMessages) packet.Payload { +func RequestMessages(ctx context.Context, sess *session.Session, request *packet.RequestMessages) packet.Payload { queries := data.New(db) - messages, err := queries.GetFrequencyMessages(ctx, request.FrequencyID) + var messages []data.Message + var err error + + if request.FrequencyID != nil && request.ReceiverID == nil { + messages, err = queries.GetFrequencyMessages(ctx, request.FrequencyID) + } else if request.ReceiverID != nil && request.FrequencyID == nil { + messages, err = queries.GetDirectMessages(ctx, data.GetDirectMessagesParams{ + User1: sess.ID(), + User2: request.ReceiverID, + }) + } else { + return &packet.Error{Error: "either receiver id or frequency id must exist"} + } + if err != nil { log.Println("database error when retrieving messages:", err) return &packet.Error{Error: "internal server error"} @@ -56,19 +63,6 @@ func GetMessages(ctx context.Context, request *packet.RequestMessages) packet.Pa return &packet.MessagesInfo{Messages: messages} } -// func GetUserById(ctx context.Context, request *packet.GetUserByID) packet.Payload { -// queries := data.New(db) -// user, err := queries.GetUserById(ctx, request.UserID) -// if err == sql.ErrNoRows { -// return &packet.Users{Users: []data.User{}} -// } -// if err != nil { -// log.Println("database error when retrieving user by id:", err) -// return &packet.ErrorMessage{Error: "internal server error"} -// } -// return &packet.Users{Users: []data.User{user}} -// } - func CreateOrGetUser(ctx context.Context, node *snowflake.Node, pubKey ed25519.PublicKey) (data.User, error) { queries := data.New(db) user, err := queries.GetUserByPublicKey(ctx, pubKey) diff --git a/internal/server/server.go b/internal/server/server.go index 5eab798..216ded2 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -157,7 +157,6 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) { } sess := session.NewSession(server, addr, user.ID, pubKey) server.AddSession(sess) - ctx = session.NewContext(ctx, sess) framer := packet.NewFramer() // Write ID back, it's useful for the client to know, and signals successful authentication @@ -200,7 +199,7 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) { if !ok { break } - response := processPacket(ctx, request) + response := processPacket(ctx, sess, request) if sess.WriteQueue == nil { break } @@ -283,36 +282,31 @@ func handleAuth(conn net.Conn, nonce []byte) (ed25519.PublicKey, error) { return pubKey, nil } -func processPacket(ctx context.Context, pkt packet.Packet) packet.Packet { - session, ok := session.FromContext(ctx) - assert.Assert(ok, "context in process packet should always have a session") - +func processPacket(ctx context.Context, sess *session.Session, pkt packet.Packet) packet.Packet { var response packet.Payload request, err := pkt.DecodedPayload() if err != nil { response = &packet.Error{Error: "malformed payload"} } else { - response = processRequest(ctx, request) + response = processRequest(ctx, sess, request) } assert.NotNil(response, "response must always be assigned to") - log.Println(session.Addr(), "sending", response.Type(), "response:", response) + log.Println(sess.Addr(), "sending", response.Type(), "response:", response) return packet.NewPacket(packet.NewMsgPackEncoder(response)) } -func processRequest(ctx context.Context, request packet.Payload) packet.Payload { - session, ok := session.FromContext(ctx) - assert.Assert(ok, "context in process packet should always have a session") - log.Println(session.Addr(), "processing", request.Type(), "request:", request) +func processRequest(ctx context.Context, sess *session.Session, request packet.Payload) packet.Payload { + log.Println(sess.Addr(), "processing", request.Type(), "request:", request) // TODO: add a way to measure the time each request/response took and log it // Potentially even separate time for code vs DB operations switch request := request.(type) { case *packet.SendMessage: - return timeout(20*time.Millisecond, api.SendMessage, ctx, request) + return timeout(20*time.Millisecond, api.SendMessage, ctx, sess, request) case *packet.RequestMessages: - return timeout(50*time.Millisecond, api.GetMessages, ctx, request) + return timeout(50*time.Millisecond, api.RequestMessages, ctx, sess, request) default: return &packet.Error{Error: "use of disallowed packet type for request"} } @@ -320,15 +314,15 @@ func processRequest(ctx context.Context, request packet.Payload) packet.Payload func timeout[T packet.Payload]( timeoutDuration time.Duration, - apiRequest func(context.Context, T) packet.Payload, - ctx context.Context, request T, + apiRequest func(context.Context, *session.Session, T) packet.Payload, + ctx context.Context, sess *session.Session, request T, ) packet.Payload { responseChan := make(chan packet.Payload) ctx, cancel := context.WithTimeout(ctx, timeoutDuration) defer cancel() go func() { - responseChan <- apiRequest(ctx, request) + responseChan <- apiRequest(ctx, sess, request) }() select { -- cgit v1.3.1