diff options
| author | Kyren223 <Kyren223@proton.me> | 2025-01-06 17:34:47 +0200 |
|---|---|---|
| committer | Kyren223 <Kyren223@proton.me> | 2025-01-06 17:34:47 +0200 |
| commit | 1a218730e9a0c16fd71f4606adfde5862be2e7b7 (patch) | |
| tree | bebf7d03f96f95174afbc0268af58dc7ca09c958 /internal/server/server.go | |
| parent | db57a31b08c1a3b00c2e253f02e67119e3cfa108 (diff) | |
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
Diffstat (limited to 'internal/server/server.go')
| -rw-r--r-- | internal/server/server.go | 39 |
1 files changed, 22 insertions, 17 deletions
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 } } |
