diff options
| author | Kyren223 <Kyren223@proton.me> | 2024-11-27 10:38:28 +0200 |
|---|---|---|
| committer | Kyren223 <Kyren223@proton.me> | 2024-11-27 10:38:28 +0200 |
| commit | f5576211d38cf8923a0e5c20ccb0677317c63ebe (patch) | |
| tree | 08eec3a29b7ed639fe8af3fc7f468d8ee1e87ea6 /internal/server | |
| parent | 328a929a8868c9ce9613feea265c64a5369709a4 (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.go | 40 | ||||
| -rw-r--r-- | internal/server/server.go | 18 |
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"} } } |
