From 94ccbc282baf45d0f2e19ab127b38c71668b650f Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Sat, 5 Jul 2025 19:50:13 +0300 Subject: Refactored server connection handling, still missing TOS and auth handling --- cmd/server/main.go | 4 +- internal/packet/packet.go | 6 +- internal/server/api/api.go | 2 +- internal/server/ctxkeys/ctxkeys.go | 8 +- internal/server/server.go | 155 +++++++++++++++++-------------------- internal/server/session/session.go | 85 +++++++++++--------- pkg/snowflake/snowflake.go | 2 + terminology.md | 6 ++ 8 files changed, 143 insertions(+), 125 deletions(-) diff --git a/cmd/server/main.go b/cmd/server/main.go index 283c8da..ac2b91d 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -46,9 +46,7 @@ func main() { }() server := server.NewServer(ctx, port) - if err := server.Run(); err != nil { - // slog.Error(err) - } + server.Run() // blocks } func setupLogging() { 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 - } - - response := processPacket(ctx, sess, request) - if ok := sess.Write(ctx, response); !ok { - return - } - } - } + writerWg.Wait() + sess.CloseWriteQueue() // causes writer to return }() - if ok := server.sendInitialPackets(ctx, sess); !ok { - return // closes the connection - } + // Processor + writerWg.Add(1) + go func() { + defer writerWg.Done() + localCtx := context.WithoutCancel(ctx) - // Infinite read loop + 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") + } + }() + + // 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() } diff --git a/pkg/snowflake/snowflake.go b/pkg/snowflake/snowflake.go index 6a53760..485026b 100644 --- a/pkg/snowflake/snowflake.go +++ b/pkg/snowflake/snowflake.go @@ -24,6 +24,8 @@ const ( type ID int64 +const InvalidID = ID(0) + func (id ID) String() string { return strconv.FormatInt(int64(id), 10) } diff --git a/terminology.md b/terminology.md index b4332c2..d4cbe91 100644 --- a/terminology.md +++ b/terminology.md @@ -15,3 +15,9 @@ Frequency ID Chat ID - A snowflake ID that refers to either a frequency ID or a receiver ID Signal - A currently open chat between 2 individual users ("DM") + +Session - A struct that helps with managing a single active connection, it contains the addr of the connection and has methods for reading/writing to the connection + +Unauthenticated Session - A session that has not yet been authenticated, so it doesn't have a User ID or public key yet, and has not been added to the server sessions map + +Authenticated Session - A session that is authenticated and attached to a user, it has a UserID and a public key and it has been mapped into the server's session map based on it's user ID. -- cgit v1.3.1