summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/server/api/api.go42
-rw-r--r--internal/server/server.go28
2 files changed, 29 insertions, 41 deletions
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 {