From 1a218730e9a0c16fd71f4606adfde5862be2e7b7 Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Mon, 6 Jan 2025 17:34:47 +0200 Subject: Server now disconnects connection when a new connection matches the same user ID, note that there seems to be a weird thing where if 2 clients are from the same IP, the disconnected client will just reconnect to the new client's socket, so they will both use the same connection for communication --- internal/server/api/helpers.go | 2 +- internal/server/server.go | 39 +++++++++++++++++++++----------------- internal/server/session/session.go | 30 +++++++++++++++++++++++++---- 3 files changed, 49 insertions(+), 22 deletions(-) (limited to 'internal') diff --git a/internal/server/api/helpers.go b/internal/server/api/helpers.go index 05dae0e..0ab6ae5 100644 --- a/internal/server/api/helpers.go +++ b/internal/server/api/helpers.go @@ -78,7 +78,7 @@ func NetworkPropagate( context, cancel := context.WithTimeout(context.Background(), timeout) go func() { defer cancel() - pkt := packet.NewPacket(packet.NewJsonEncoder(payload)) + pkt := packet.NewPacket(packet.NewMsgPackEncoder(payload)) if ok := session.Write(context, pkt); !ok { log.Println(sess.Addr(), "propagation to", session.Addr(), "failed") } diff --git a/internal/server/server.go b/internal/server/server.go index 0534023..495b2a5 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -71,6 +71,9 @@ func NewServer(ctx context.Context, port uint16) server { func (s *server) AddSession(session *session.Session) { s.sessMu.Lock() defer s.sessMu.Unlock() + if sess, ok := s.sessions[session.ID()]; ok { + sess.Close() + } s.sessions[session.ID()] = session } @@ -142,7 +145,7 @@ func (server *server) handleConnection(conn net.Conn) { log.Println(addr, "accepted") - initialCtx, cancel := context.WithTimeout(server.ctx, 5*time.Second) + initialCtx, initialCancel := context.WithTimeout(server.ctx, 5*time.Second) deadline, _ := initialCtx.Deadline() err := conn.SetDeadline(deadline) assert.NoError(err, "setting read deadline should not error") @@ -152,7 +155,7 @@ func (server *server) handleConnection(conn net.Conn) { pubKey, err := handleAuth(conn) if err != nil { - cancel() + initialCancel() log.Println(addr, err) conn.Close() log.Println(addr, "disconnected") @@ -161,13 +164,15 @@ func (server *server) handleConnection(conn net.Conn) { user, err := api.CreateOrGetUser(initialCtx, server.Node(), pubKey) if err != nil { - cancel() + initialCancel() log.Println(addr, "user creation/fetching error:", err) conn.Close() log.Println(addr, "disconnected") return } - sess := session.NewSession(server, addr, user.ID, pubKey) + ctx, cancel := context.WithCancel(server.ctx) + defer cancel() + sess := session.NewSession(server, addr, cancel, user.ID, pubKey) server.AddSession(sess) framer := packet.NewFramer() @@ -176,17 +181,17 @@ func (server *server) handleConnection(conn net.Conn) { binary.BigEndian.PutUint64(id[:], uint64(user.ID)) _, err = conn.Write(id[:]) if err != nil { - cancel() + initialCancel() log.Println(addr, "failed to write user id") conn.Close() log.Println(addr, "disconnected") return } - cancel() + initialCancel() go func() { - <-server.ctx.Done() + <-ctx.Done() conn.Close() }() defer func() { @@ -197,7 +202,7 @@ func (server *server) handleConnection(conn net.Conn) { go func() { for { - packet, ok := sess.Read(server.ctx) + packet, ok := sess.Read(ctx) if !ok { return } @@ -212,15 +217,15 @@ func (server *server) handleConnection(conn net.Conn) { go func() { for { select { - case <-server.ctx.Done(): + case <-ctx.Done(): return case request, ok := <-framer.Out: if !ok { return } - response := processPacket(server.ctx, sess, request) - if ok := sess.Write(server.ctx, response); !ok { + response := processPacket(ctx, sess, request) + if ok := sess.Write(ctx, response); !ok { return } } @@ -228,12 +233,12 @@ func (server *server) handleConnection(conn net.Conn) { }() // Send initial packets - payload, err := api.GetNetworksInfo(server.ctx, sess) + payload, err := api.GetNetworksInfo(ctx, sess) if err != nil { return // closes the connection } infoPacket := packet.NewPacket(packet.NewMsgPackEncoder(payload)) - sess.Write(server.ctx, infoPacket) + sess.Write(ctx, infoPacket) buffer := make([]byte, 512) for { @@ -245,15 +250,15 @@ func (server *server) handleConnection(conn net.Conn) { break } - err = framer.Push(server.ctx, buffer[:n]) - if server.ctx.Err() != nil { - log.Println(addr, server.ctx.Err()) + err = framer.Push(ctx, buffer[:n]) + if ctx.Err() != nil { + log.Println(addr, ctx.Err()) break } if err != nil { payload := packet.Error{Error: err.Error()} pkt := packet.NewPacket(packet.NewMsgPackEncoder(&payload)) - sess.Write(server.ctx, pkt) + sess.Write(ctx, pkt) break } } diff --git a/internal/server/session/session.go b/internal/server/session/session.go index 2b7121a..88a19e7 100644 --- a/internal/server/session/session.go +++ b/internal/server/session/session.go @@ -25,6 +25,7 @@ type SessionManager interface { type Session struct { manager SessionManager addr *net.TCPAddr + cancel context.CancelFunc writeQueue chan packet.Packet issuedTime time.Time @@ -36,14 +37,21 @@ type Session struct { mu sync.Mutex } -func NewSession(manager SessionManager, addr *net.TCPAddr, id snowflake.ID, pubKey ed25519.PublicKey) *Session { +func NewSession( + manager SessionManager, + addr *net.TCPAddr, cancel context.CancelFunc, + id snowflake.ID, pubKey ed25519.PublicKey, +) *Session { session := &Session{ - writeQueue: make(chan packet.Packet, 10), - PubKey: pubKey, manager: manager, addr: addr, - id: id, + cancel: cancel, + writeQueue: make(chan packet.Packet, 10), + issuedTime: time.Time{}, challenge: make([]byte, 32), + PubKey: pubKey, + id: id, + mu: sync.Mutex{}, } session.Challenge() // Make sure an initial nonce is generated return session @@ -89,3 +97,17 @@ func (s *Session) Read(ctx context.Context) (packet.Packet, bool) { return packet.Packet{}, false } } + +func (s *Session) Close() { + timeout := 1 * time.Second + ctx, cancel := context.WithTimeout(context.Background(), timeout) + payload := &packet.Error{ + Error: "new connection from another location, closing this one", + PktType: packet.PacketError, + } + pkt := packet.NewPacket(packet.NewMsgPackEncoder(payload)) + s.Write(ctx, pkt) + cancel() + + s.cancel() +} -- cgit v1.3.1