diff --git a/internal/common/packets.go b/internal/common/packets.go index f591538..ec53c93 100644 --- a/internal/common/packets.go +++ b/internal/common/packets.go @@ -6,6 +6,11 @@ import ( "tangled.org/kutuptilkisi/RIRC/internal/utils" ) +const ( + MAX_MESSAGE_LENGTH = 1024 + MAX_USERNAME_LENGTH = 32 +) + type PacketFactory = func() Packet var ServerboundPackets map[uint16]PacketFactory = map[uint16]PacketFactory{ @@ -33,7 +38,7 @@ type AuthenticatePacket struct { } func (p *AuthenticatePacket) Decode(r io.Reader) error { - name, err := utils.ReadString(r) + name, err := utils.ReadStringLimit(r, MAX_MESSAGE_LENGTH) if err != nil { return err } @@ -51,7 +56,7 @@ type ServerboundMessagePacket struct { } func (p *ServerboundMessagePacket) Decode(r io.Reader) error { - message, err := utils.ReadString(r) + message, err := utils.ReadStringLimit(r, MAX_MESSAGE_LENGTH) if err != nil { return err } @@ -69,13 +74,13 @@ type ClientboundMessagePacket struct { } func (p *ClientboundMessagePacket) Decode(r io.Reader) error { - message, err := utils.ReadString(r) + message, err := utils.ReadStringLimit(r, MAX_MESSAGE_LENGTH) if err != nil { return err } p.Message = message - sender, err := utils.ReadString(r) + sender, err := utils.ReadStringLimit(r, MAX_USERNAME_LENGTH) if err != nil { return err } diff --git a/internal/server/server.go b/internal/server/server.go index 99a6737..27133d6 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -32,12 +32,15 @@ func (s *Server) StartServer(port uint16) error { } } -const MAX_PACKET_LENGTH = 1024 * 32 +const ( + MAX_PACKET_LENGTH = 1024 * 32 + MAX_BUFFER_SIZE = 1024 * 8 +) func handleConnection(s *Server, conn net.Conn) { defer conn.Close() - reader := bufio.NewReaderSize(conn, 4096*2) - header := make([]byte, 6) + reader := bufio.NewReaderSize(conn, MAX_BUFFER_SIZE) + header := make([]byte, 6) // 4 byte length + 2 byte packet id user := User{ id: uuid.New(), @@ -47,6 +50,8 @@ func handleConnection(s *Server, conn net.Conn) { s.AddUser(&user) for { + // NOTE: ReadFull might be growing the header buffer. + // In that case we should fix it if _, err := io.ReadFull(reader, header); err != nil { if err != io.EOF { log.Println(err) @@ -64,7 +69,9 @@ func handleConnection(s *Server, conn net.Conn) { user.conn = nil return } + packetId := binary.BigEndian.Uint16(header[4:]) + packetReader := io.LimitReader(reader, int64(length)) handlePacket(packetReader, &user, packetId) } diff --git a/internal/utils/string.go b/internal/utils/string.go index fca6401..a247949 100644 --- a/internal/utils/string.go +++ b/internal/utils/string.go @@ -2,15 +2,22 @@ package utils import ( "encoding/binary" + "errors" "io" ) -func ReadString(r io.Reader) (string, error) { +func ReadStringLimit(r io.Reader, limit uint16) (string, error) { var length uint16 if err := binary.Read(r, binary.BigEndian, &length); err != nil { return "", err } + if length > limit { + // TODO: Check properly + // TODO: Probably errors should have their own package + return "", errors.New("Passed Limit") + } + if length == 0 { return "", nil }