summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/packet/packet.go6
-rw-r--r--internal/server/api/api.go2
-rw-r--r--internal/server/ctxkeys/ctxkeys.go8
-rw-r--r--internal/server/server.go153
-rw-r--r--internal/server/session/session.go85
5 files changed, 133 insertions, 121 deletions
diff --git a/internal/packet/packet.go b/internal/packet/packet.go
index 543b227..d3dcd7e 100644
--- a/internal/packet/packet.go
+++ b/internal/packet/packet.go
@@ -169,7 +169,7 @@ func (e PacketType) String() string {
}
const (
- VERSION = byte(1)
+ VERSION = byte(2)
PACKET_MAX_SIZE = math.MaxUint16
PAYLOAD_MAX_SIZE = PACKET_MAX_SIZE - HEADER_SIZE
HEADER_SIZE = 4
@@ -355,9 +355,11 @@ type PacketFramer struct {
buffer []byte
}
+const ReadQueueSize = 10
+
func NewFramer() PacketFramer {
return PacketFramer{
- Out: make(chan Packet, 10),
+ Out: make(chan Packet, ReadQueueSize),
}
}
diff --git a/internal/server/api/api.go b/internal/server/api/api.go
index c0bf3cb..dddd047 100644
--- a/internal/server/api/api.go
+++ b/internal/server/api/api.go
@@ -167,7 +167,7 @@ func SendMessage(ctx context.Context, sess *session.Session, request *packet.Sen
log.Println("api.go:167 database error:", err)
return &ErrInternalError
}
- if !bytes.Equal(sess.PubKey, pubKey) {
+ if !bytes.Equal(sess.PubKey(), pubKey) {
return &ErrPermissionDenied
}
}
diff --git a/internal/server/ctxkeys/ctxkeys.go b/internal/server/ctxkeys/ctxkeys.go
index c0f688d..3677a07 100644
--- a/internal/server/ctxkeys/ctxkeys.go
+++ b/internal/server/ctxkeys/ctxkeys.go
@@ -13,11 +13,15 @@ const (
UserID key = iota
IpAddr
KeyMax
+ Evicted
+ EvictedBy
)
var keyNames = map[key]string{
- UserID: "user_id",
- IpAddr: "ip_addr",
+ UserID: "user_id",
+ IpAddr: "ip_addr",
+ Evicted: "evicted",
+ EvictedBy: "evicted_by",
}
func Init() {
diff --git a/internal/server/server.go b/internal/server/server.go
index adcde77..5791f8e 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -5,7 +5,6 @@ import (
"crypto/ed25519"
"crypto/rand"
"crypto/tls"
- "encoding/binary"
"errors"
"fmt"
"io"
@@ -75,15 +74,39 @@ func NewServer(ctx context.Context, port uint16) server {
}
}
-func (s *server) AddSession(session *session.Session) {
+func (s *server) AddSession(session *session.Session, userId snowflake.ID, pubKey ed25519.PublicKey) {
s.sessMu.Lock()
defer s.sessMu.Unlock()
+
+ session.Promote(userId, pubKey)
+
if sess, ok := s.sessions[session.ID()]; ok {
- sess.Close()
+ EvictSession(sess) // last connection wins
+ slog.Info("closed due to new connection from another location",
+ ctxkeys.IpAddr, sess.Addr(), ctxkeys.UserID, sess.ID(),
+ ctxkeys.EvictedBy, session.Addr())
+ slog.Info("this session evicted another session",
+ ctxkeys.IpAddr, session.Addr(), ctxkeys.UserID, session.ID(),
+ ctxkeys.Evicted, sess.Addr())
}
+
s.sessions[session.ID()] = session
}
+func EvictSession(sess *session.Session) {
+ timeout := 10 * time.Millisecond
+ ctx, cancel := context.WithTimeout(context.Background(), timeout)
+ defer cancel()
+ payload := &packet.Error{
+ Error: "new connection from another location, closing this one",
+ PktType: packet.PacketError,
+ }
+ pkt := packet.NewPacket(packet.NewMsgPackEncoder(payload))
+ sess.Write(ctx, pkt)
+
+ sess.Close()
+}
+
func (s *server) RemoveSession(id snowflake.ID) {
s.sessMu.Lock()
defer s.sessMu.Unlock()
@@ -109,7 +132,7 @@ func (s *server) Node() *snowflake.Node {
// Run starts listening and accepting clients,
// blocking until it gets terminated by cancelling the context.
-func (s *server) Run() error {
+func (s *server) Run() {
listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(int(s.Port)), tlsConfig)
if err != nil {
log.Fatalf("error starting server: %s", err)
@@ -132,7 +155,7 @@ func (s *server) Run() error {
}
if s.ctx.Err() != nil {
- slog.Info("server context expired", "error", s.ctx.Err())
+ slog.Info("server context done", "error", s.ctx.Err())
break
}
continue // Ignore and skip (don't connect)
@@ -148,7 +171,6 @@ func (s *server) Run() error {
slog.Info("waiting for all active connections to close...")
wg.Wait()
slog.Info("completed server shutdown")
- return nil
}
func (server *server) handleConnection(conn net.Conn) {
@@ -162,126 +184,95 @@ func (server *server) handleConnection(conn net.Conn) {
slog.InfoContext(ctx, "connection accepted")
defer slog.InfoContext(ctx, "connection closed")
+ defer conn.Close()
- // Set deadline before auth
- deadline := time.Now().Add(5 * time.Second)
- err := conn.SetDeadline(deadline)
- assert.NoError(err, "setting deadline should not error")
-
- pubKey, err := handleAuth(conn)
- if err != nil {
- slog.Info("user authentication failed", "error", err)
- _ = conn.Close()
- return
- }
-
- // Reset deadline after auth
- err = conn.SetDeadline(time.Time{})
- assert.NoError(err, "unsetting deadline should not error")
-
- user, err := api.CreateOrGetUser(ctx, server.Node(), pubKey)
- if err != nil {
- log.Println(addr, "user creation/fetching error:", err)
- _ = conn.Close()
- return
- }
-
- ctx = context.WithValue(ctx, ctxkeys.UserID, user.ID)
- sess := session.NewSession(server, addr, cancel, user.ID, pubKey)
- server.AddSession(sess)
+ var writerWg *sync.WaitGroup
+ done := make(chan struct{})
framer := packet.NewFramer()
- // Write ID back, it's useful for the client to know, and signals successful authentication
- var id [8]byte
- binary.BigEndian.PutUint64(id[:], uint64(user.ID)) // #nosec G115 -- sign bit is always 0 in snowflake IDs
- _, err = conn.Write(id[:])
- if err != nil {
- log.Println(addr, "failed to write user id")
- _ = conn.Close()
- log.Println(addr, "disconnected")
- return
- }
-
+ sess := session.NewSession(server, addr, cancel, writerWg)
go func() {
<-ctx.Done()
- _ = conn.Close()
- }()
- defer func() {
- _ = conn.Close()
- sameAddress := addr.String() == server.Session(sess.ID()).Addr().String()
- if sameAddress {
- server.RemoveSession(sess.ID())
+ // Remove session after cancellation
+ if sess.IsAuthenticated() {
+ sameAddress := addr.String() == server.Session(sess.ID()).Addr().String()
+ // false if the user signed in from a different connection
+ if sameAddress {
+ server.RemoveSession(sess.ID())
+ }
}
- log.Println(addr, "disconnected gracefully")
}()
+ // Writer
go func() {
- for {
- packet, ok := sess.Read(ctx)
- if !ok {
- return
- }
- log.Println(addr, "sending packet:", packet)
+ defer close(done)
+ writeQueue := sess.Read()
+
+ for packet := range writeQueue {
if _, err := packet.Into(conn); err != nil {
// TODO: probably should add this to prevent the
// "use of closed connection" error, as it's intended to happen
+ // and once it happens we can just return
// if !errors.Is(err, net.ErrClosed) {
// log.Println(addr, err)
// }
- log.Println(addr, err)
+ slog.ErrorContext(ctx, "error sending packet", "error", err, "packet", packet)
return
}
+ slog.InfoContext(ctx, "packet sent", "packet", packet)
}
}()
+ // Writer closer
go func() {
- for {
- select {
- case <-ctx.Done():
- return
- case request, ok := <-framer.Out:
- if !ok {
- return
- }
+ writerWg.Wait()
+ sess.CloseWriteQueue() // causes writer to return
+ }()
- response := processPacket(ctx, sess, request)
- if ok := sess.Write(ctx, response); !ok {
- return
- }
- }
+ // Processor
+ writerWg.Add(1)
+ go func() {
+ defer writerWg.Done()
+ localCtx := context.WithoutCancel(ctx)
+
+ for request := range framer.Out {
+ response := processPacket(localCtx, sess, request)
+ ok = sess.Write(localCtx, response)
+ assert.Assert(ok, "context is never done and write will panic")
}
}()
- if ok := server.sendInitialPackets(ctx, sess); !ok {
- return // closes the connection
- }
-
- // Infinite read loop
+ // Reader
buffer := make([]byte, 512)
for {
n, err := conn.Read(buffer)
if err != nil {
- if !errors.Is(err, io.EOF) {
- log.Println(addr, err)
+ if errors.Is(err, io.EOF) {
+ slog.InfoContext(ctx, "closed gracefully")
} else {
- log.Println(addr, "disconnecting gracefully...")
+ slog.ErrorContext(ctx, "failed reading from buffer", "error", err)
}
break
}
err = framer.Push(ctx, buffer[:n])
if ctx.Err() != nil {
- log.Println(addr, ctx.Err())
+ slog.InfoContext(ctx, "reader context done", "error", ctx.Err())
break
}
if err != nil {
payload := packet.Error{Error: err.Error()}
pkt := packet.NewPacket(packet.NewMsgPackEncoder(&payload))
+ writerWg.Add(1)
sess.Write(ctx, pkt)
- log.Println(addr, "received malformed packet:", err)
+ writerWg.Done()
+ slog.WarnContext(ctx, "received malformed packet", "error", err)
break
}
}
+ close(framer.Out) // stop processing
+
+ <-done
}
func handleAuth(conn net.Conn) (ed25519.PublicKey, error) {
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()
}