summaryrefslogtreecommitdiff
path: root/internal/server/session/session.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/server/session/session.go')
-rw-r--r--internal/server/session/session.go85
1 files changed, 50 insertions, 35 deletions
diff --git a/internal/server/session/session.go b/internal/server/session/session.go
index a904d39..ffe4e13 100644
--- a/internal/server/session/session.go
+++ b/internal/server/session/session.go
@@ -4,7 +4,6 @@ import (
"context"
"crypto/ed25519"
"crypto/rand"
- "log"
"net"
"sync"
"time"
@@ -14,8 +13,10 @@ import (
"github.com/kyren223/eko/pkg/snowflake"
)
+const WriteQueueSize = 10
+
type SessionManager interface {
- AddSession(session *Session)
+ AddSession(session *Session, userId snowflake.ID, pubKey ed25519.PublicKey)
RemoveSession(id snowflake.ID)
Session(id snowflake.ID) *Session
UseSessions(f func(map[snowflake.ID]*Session))
@@ -24,15 +25,18 @@ type SessionManager interface {
}
type Session struct {
- manager SessionManager
- addr *net.TCPAddr
- cancel context.CancelFunc
+ manager SessionManager
+ addr *net.TCPAddr
+ cancel context.CancelFunc
+
writeQueue chan packet.Packet
+ writerWg *sync.WaitGroup
+ writeMu sync.RWMutex
issuedTime time.Time
challenge []byte
- PubKey ed25519.PublicKey
+ pubKey ed25519.PublicKey
id snowflake.ID
mu sync.Mutex
@@ -41,31 +45,49 @@ type Session struct {
func NewSession(
manager SessionManager,
addr *net.TCPAddr, cancel context.CancelFunc,
- id snowflake.ID, pubKey ed25519.PublicKey,
+ writerWg *sync.WaitGroup,
) *Session {
+ assert.NotNil(addr, "tcp address should be valid")
+ assert.NotNil(manager, "session manager should be valid")
session := &Session{
manager: manager,
addr: addr,
cancel: cancel,
- writeQueue: make(chan packet.Packet, 10),
+ writeQueue: make(chan packet.Packet, WriteQueueSize),
+ writerWg: writerWg,
issuedTime: time.Time{},
challenge: make([]byte, 32),
- PubKey: pubKey,
- id: id,
+ pubKey: ed25519.PublicKey{},
+ id: snowflake.InvalidID,
mu: sync.Mutex{},
}
- session.Challenge() // Make sure an initial nonce is generated
return session
}
func (s *Session) Addr() *net.TCPAddr {
+ assert.NotNil(s.addr, "tcp address should be valid")
return s.addr
}
+func (s *Session) IsAuthenticated() bool {
+ return s.id != snowflake.InvalidID
+}
+
func (s *Session) ID() snowflake.ID {
+ assert.Assert(s.IsAuthenticated(), "use of ID in an unauthenticated session", "addr", s.addr)
return s.id
}
+func (s *Session) PubKey() ed25519.PublicKey {
+ assert.Assert(s.IsAuthenticated(), "use of PubKey in an unauthenticated session", "addr", s.addr)
+ return s.pubKey
+}
+
+func (s *Session) Promote(userId snowflake.ID, pubKey ed25519.PublicKey) {
+ s.id = userId
+ s.pubKey = pubKey
+}
+
func (s *Session) Manager() SessionManager {
return s.manager
}
@@ -82,6 +104,12 @@ func (s *Session) Challenge() []byte {
}
func (s *Session) Write(ctx context.Context, pkt packet.Packet) bool {
+ s.writerWg.Add(1)
+ defer s.writerWg.Done()
+
+ s.writeMu.RLock()
+ defer s.writeMu.RUnlock()
+
select {
case s.writeQueue <- pkt:
return true
@@ -90,33 +118,20 @@ func (s *Session) Write(ctx context.Context, pkt packet.Packet) bool {
}
}
-func (s *Session) Read(ctx context.Context) (packet.Packet, bool) {
- select {
- case pkt := <-s.writeQueue:
- return pkt, true
- case <-ctx.Done():
- return packet.Packet{}, false
- }
+func (s *Session) Read() <-chan packet.Packet {
+ s.writeMu.RLock()
+ defer s.writeMu.RUnlock()
+ return s.writeQueue
}
-func (s *Session) Close() {
- timeout := 10 * time.Millisecond
- 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()
-
- // Add some delay before canceling to let the writer enough time to
- // actually write that into the connection
- // HACK: consider just giving session the private connection and to
- // write directly so we don't have to wait
- time.Sleep(100 * time.Millisecond)
+func (s *Session) CloseWriteQueue() {
+ s.writeMu.Lock()
+ defer s.writeMu.Unlock()
- log.Println(s.addr, "closed due to new connection from another location")
+ close(s.writeQueue)
+ s.writeQueue = nil
+}
+func (s *Session) Close() {
s.cancel()
}