summaryrefslogtreecommitdiff
path: root/internal/server
diff options
context:
space:
mode:
authorKyren223 <Kyren223@proton.me>2025-01-06 17:34:47 +0200
committerKyren223 <Kyren223@proton.me>2025-01-06 17:34:47 +0200
commit1a218730e9a0c16fd71f4606adfde5862be2e7b7 (patch)
treebebf7d03f96f95174afbc0268af58dc7ca09c958 /internal/server
parentdb57a31b08c1a3b00c2e253f02e67119e3cfa108 (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')
-rw-r--r--internal/server/api/helpers.go2
-rw-r--r--internal/server/server.go39
-rw-r--r--internal/server/session/session.go30
3 files changed, 49 insertions, 22 deletions
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()
+}