summaryrefslogtreecommitdiff
path: root/internal/server
diff options
context:
space:
mode:
authorKyren223 <Kyren223@proton.me>2024-11-27 10:38:28 +0200
committerKyren223 <Kyren223@proton.me>2024-11-27 10:38:28 +0200
commitf5576211d38cf8923a0e5c20ccb0677317c63ebe (patch)
tree08eec3a29b7ed639fe8af3fc7f468d8ee1e87ea6 /internal/server
parent328a929a8868c9ce9613feea265c64a5369709a4 (diff)
Fixed server-side issues after refactoring, more work needs to be done
to refactor server
Diffstat (limited to 'internal/server')
-rw-r--r--internal/server/api/api.go40
-rw-r--r--internal/server/server.go18
2 files changed, 27 insertions, 31 deletions
diff --git a/internal/server/api/api.go b/internal/server/api/api.go
index 1dc74d2..f60e0ec 100644
--- a/internal/server/api/api.go
+++ b/internal/server/api/api.go
@@ -20,12 +20,12 @@ func SendMessage(ctx context.Context, request *packet.SendMessage) packet.Payloa
assert.Assert(ok, "context in process packet should always have a session")
if (request.ReceiverID != nil) == (request.FrequencyID != nil) {
- return &packet.ErrorMessage{Error: "either receiver id or frequency id must exist"}
+ return &packet.Error{Error: "either receiver id or frequency id must exist"}
}
content := strings.TrimSpace(request.Content)
if content == "" {
- return &packet.ErrorMessage{Error: "message content must not be blank"}
+ return &packet.Error{Error: "message content must not be blank"}
}
node := sess.Manager().Node()
@@ -40,34 +40,34 @@ func SendMessage(ctx context.Context, request *packet.SendMessage) packet.Payloa
})
if err != nil {
log.Println(sess.Addr(), "SendMessage database error:", err)
- return &packet.ErrorMessage{Error: "internal server error"}
+ return &packet.Error{Error: "internal server error"}
}
- return &packet.Messages{Messages: []data.Message{message}}
+ return &packet.MessagesInfo{Messages: []data.Message{message}}
}
-func GetMessages(ctx context.Context, request *packet.GetMessagesRange) packet.Payload {
+func GetMessages(ctx context.Context, request *packet.RequestMessages) packet.Payload {
queries := data.New(db)
messages, err := queries.GetFrequencyMessages(ctx, request.FrequencyID)
if err != nil {
log.Println("database error when retrieving messages:", err)
- return &packet.ErrorMessage{Error: "internal server error"}
+ return &packet.Error{Error: "internal server error"}
}
- return &packet.Messages{Messages: messages}
+ 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 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)
@@ -76,7 +76,7 @@ func CreateOrGetUser(ctx context.Context, node *snowflake.Node, pubKey ed25519.P
id := node.Generate()
user, err = queries.CreateUser(ctx, data.CreateUserParams{
ID: id,
- Name: "User" + strconv.FormatInt(id.Time() % 1000, 10),
+ Name: "User" + strconv.FormatInt(id.Time()%1000, 10),
PublicKey: pubKey,
})
}
diff --git a/internal/server/server.go b/internal/server/server.go
index dce8938..5eab798 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -186,9 +186,7 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
if !ok {
break
}
- if packet.Type().IsPush() {
- log.Println(addr, "streaming packet:", packet)
- }
+ log.Println(addr, "sending packet:", packet)
if _, err := packet.Into(conn); err != nil {
log.Println(addr, err)
break
@@ -212,6 +210,7 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
buffer := make([]byte, 512)
for {
+ // TODO: do we need this deadline
err := conn.SetReadDeadline(time.Now().Add(time.Second))
assert.NoError(err, "setting read deadline should not error")
n, err := conn.Read(buffer)
@@ -233,7 +232,7 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
if ctx.Err() != nil {
log.Println(addr, ctx.Err())
} else {
- payload := packet.ErrorMessage{Error: err.Error()}
+ payload := packet.Error{Error: err.Error()}
pkt := packet.NewPacket(packet.NewMsgPackEncoder(&payload))
sess.WriteQueue <- pkt
}
@@ -292,7 +291,7 @@ func processPacket(ctx context.Context, pkt packet.Packet) packet.Packet {
request, err := pkt.DecodedPayload()
if err != nil {
- response = &packet.ErrorMessage{Error: "malformed payload"}
+ response = &packet.Error{Error: "malformed payload"}
} else {
response = processRequest(ctx, request)
}
@@ -312,12 +311,10 @@ func processRequest(ctx context.Context, request packet.Payload) packet.Payload
switch request := request.(type) {
case *packet.SendMessage:
return timeout(20*time.Millisecond, api.SendMessage, ctx, request)
- case *packet.GetMessagesRange:
+ case *packet.RequestMessages:
return timeout(50*time.Millisecond, api.GetMessages, ctx, request)
- case *packet.GetUserByID:
- return timeout(50*time.Millisecond, api.GetUserById, ctx, request)
default:
- return &packet.ErrorMessage{Error: "use of disallowed packet type for request"}
+ return &packet.Error{Error: "use of disallowed packet type for request"}
}
}
@@ -341,7 +338,6 @@ func timeout[T packet.Payload](
sess, ok := session.FromContext(ctx)
assert.Assert(ok, "session should exist")
log.Println(sess.Addr(), "timeout of", request.Type(), "request")
- // TODO: consider if we want to say it's a timeout or be vague to mitigate DOS attacks
- return &packet.ErrorMessage{Error: "internal server error"}
+ return &packet.Error{Error: "request timeout"}
}
}