summaryrefslogtreecommitdiff
path: root/internal/server
diff options
context:
space:
mode:
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()
+}