From 6a34d7be01256df472015564a655d8253acd8724 Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Wed, 27 Nov 2024 12:52:22 +0200 Subject: Refactored server infrastructure --- cmd/server/main.go | 12 ++++--- internal/packet/packet.go | 23 +++++++++---- internal/server/api/database.go | 6 ++-- internal/server/server.go | 68 +++++++++++++++----------------------- internal/server/session/session.go | 35 +++++++++----------- 5 files changed, 67 insertions(+), 77 deletions(-) diff --git a/cmd/server/main.go b/cmd/server/main.go index 0359ff8..6ac03ea 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -11,6 +11,7 @@ import ( "github.com/kyren223/eko/internal/server" "github.com/kyren223/eko/internal/server/api" + "github.com/kyren223/eko/pkg/assert" ) const port = 7223 @@ -32,11 +33,11 @@ func main() { } api.ConnectToDatabase() - defer api.CloseDatabase() - - server := server.NewServer(port) + assert.AddFlush(api.DB()) + defer api.DB().Close() ctx, cancel := context.WithCancel(context.Background()) + defer cancel() signalChan := make(chan os.Signal, 1) signal.Notify(signalChan, syscall.SIGINT, syscall.SIGTERM) go func() { @@ -45,5 +46,8 @@ func main() { cancel() }() - server.ListenAndServe(ctx) + server := server.NewServer(port) + if err := server.Run(ctx); err != nil { + log.Println(err) + } } diff --git a/internal/packet/packet.go b/internal/packet/packet.go index 8eb1a86..7476056 100644 --- a/internal/packet/packet.go +++ b/internal/packet/packet.go @@ -222,14 +222,14 @@ func (p Packet) DecodedPayload() (Payload, error) { } var ( - PacketUnsupportedVersion error = errors.New("packet error: unsupported version") - PacketUnsupportedEncoding error = errors.New("packet error: unsupported encoding") - PacketUnsupportedType error = errors.New("packet error: unsupported type") + ErrUnsupportedVersion error = errors.New("packet error: unsupported version") + ErrUnsupportedEncoding error = errors.New("packet error: unsupported encoding") + ErrUnsupportedType error = errors.New("packet error: unsupported type") ) type PacketFramer struct { - buffer []byte Out chan Packet + buffer []byte } func NewFramer() PacketFramer { @@ -239,13 +239,22 @@ func NewFramer() PacketFramer { } func (f *PacketFramer) Push(ctx context.Context, data []byte) error { + if ctx.Err() != nil { + return ctx.Err() + } + f.buffer = append(f.buffer, data...) for { + if ctx.Err() != nil { + return ctx.Err() + } + packet, err := f.parse() if packet == nil || err != nil { return err } + select { case f.Out <- *packet: case <-ctx.Done(): @@ -260,17 +269,17 @@ func (f *PacketFramer) parse() (*Packet, error) { } if f.buffer[VERSION_OFFSET] != VERSION { - return nil, PacketUnsupportedVersion + return nil, ErrUnsupportedVersion } encoding := Encoding(f.buffer[ENCODING_OFFSET] >> 6) if !encoding.IsSupported() { - return nil, PacketUnsupportedEncoding + return nil, ErrUnsupportedEncoding } packetType := PacketType(f.buffer[TYPE_OFFSET] & 63) if !packetType.IsSupported() { - return nil, PacketUnsupportedType + return nil, ErrUnsupportedType } length := binary.BigEndian.Uint16(f.buffer[LENGTH_OFFSET:]) diff --git a/internal/server/api/database.go b/internal/server/api/database.go index 3546395..a137746 100644 --- a/internal/server/api/database.go +++ b/internal/server/api/database.go @@ -36,10 +36,8 @@ func ConnectToDatabase() { log.Println("database connection ready to be used") } -func CloseDatabase() { - assert.NotNil(db, "db should only be closed if it exists") - db.Close() - log.Println("connection with database closed") +func DB() *sql.DB { + return db } func demo() { diff --git a/internal/server/server.go b/internal/server/server.go index 216ded2..5f4376a 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -12,7 +12,6 @@ import ( "io" "log" "net" - "os" "strconv" "sync" "time" @@ -48,9 +47,9 @@ func init() { type server struct { node *snowflake.Node - Port uint16 sessions map[snowflake.ID]*session.Session sessMu sync.RWMutex + Port uint16 } // Creates a new server on the given port. @@ -90,12 +89,9 @@ func (s *server) Node() *snowflake.Node { return s.node } -// Starts listening and accepting clients on the server's port. -// -// The given context is used for cancellation, -// note that the server will wait for all active connections to close before -// returning, this is a blocking operation. -func (s *server) ListenAndServe(ctx context.Context) { +// Run starts listening and accepting clients, +// blocking until it gets terminated by cancelling the context. +func (s *server) Run(ctx context.Context) error { listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(int(s.Port)), tlsConfig) if err != nil { log.Fatalf("error starting server: %s", err) @@ -108,7 +104,7 @@ func (s *server) ListenAndServe(ctx context.Context) { listener.Close() }() - log.Printf("started listening on port %v...\n", s.Port) + log.Println("started listening on port", s.Port) var wg sync.WaitGroup for { conn, err := listener.Accept() @@ -120,27 +116,25 @@ func (s *server) ListenAndServe(ctx context.Context) { } wg.Add(1) go func() { - handleConnection(ctx, conn, s) + s.handleConnection(ctx, conn) wg.Done() }() } - log.Printf("stopped listening on port %v\n", s.Port) + log.Println("stopped listening on port", s.Port) log.Println("waiting for all active connections to close...") wg.Wait() log.Println("server shutdown complete") + return nil } -func handleConnection(ctx context.Context, conn net.Conn, server *server) { +func (server *server) handleConnection(ctx context.Context, conn net.Conn) { addr, ok := conn.RemoteAddr().(*net.TCPAddr) assert.Assert(ok, "getting tcp address should never fail as we are using tcp connections") log.Println(addr, "accepted") - nonce := [32]byte{} - _, err := rand.Read(nonce[:]) - assert.NoError(err, "random should always produce a value") - pubKey, err := handleAuth(conn, nonce[:]) + pubKey, err := handleAuth(conn) if err != nil { log.Println(addr, err) conn.Close() @@ -172,10 +166,8 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) { defer func() { conn.Close() - close(framer.Out) server.RemoveSession(sess.ID()) - close(sess.WriteQueue) - sess.WriteQueue = nil + close(framer.Out) // avoid leaking goroutine log.Println(addr, "disconnected") }() @@ -200,48 +192,42 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) { break } response := processPacket(ctx, sess, request) - if sess.WriteQueue == nil { + if ok := sess.Write(ctx, response); !ok { break } - sess.WriteQueue <- response } }() buffer := make([]byte, 512) for { - // TODO: do we need this deadline - err := conn.SetReadDeadline(time.Now().Add(time.Second)) - assert.NoError(err, "setting read deadline should not error") n, err := conn.Read(buffer) - deadlineExceeded := errors.Is(err, os.ErrDeadlineExceeded) - if err != nil && !deadlineExceeded { + if err != nil { if !errors.Is(err, io.EOF) { - log.Println(addr, "read error:", err) + log.Println(addr, err) } break } + err = framer.Push(ctx, buffer[:n]) if ctx.Err() != nil { log.Println(addr, ctx.Err()) break } - - err = framer.Push(ctx, buffer[:n]) if err != nil { - if ctx.Err() != nil { - log.Println(addr, ctx.Err()) - } else { - payload := packet.Error{Error: err.Error()} - pkt := packet.NewPacket(packet.NewMsgPackEncoder(&payload)) - sess.WriteQueue <- pkt - } + payload := packet.Error{Error: err.Error()} + pkt := packet.NewPacket(packet.NewMsgPackEncoder(&payload)) + sess.Write(ctx, pkt) break } } } -func handleAuth(conn net.Conn, nonce []byte) (ed25519.PublicKey, error) { - err := conn.SetDeadline(time.Now().Add(time.Second * 5)) +func handleAuth(conn net.Conn) (ed25519.PublicKey, error) { + nonce := [32]byte{} + _, err := rand.Read(nonce[:]) + assert.NoError(err, "random should always produce a value") + + err = conn.SetDeadline(time.Now().Add(time.Second * 5)) assert.NoError(err, "setting read deadline should not error") defer func() { @@ -251,7 +237,7 @@ func handleAuth(conn net.Conn, nonce []byte) (ed25519.PublicKey, error) { challengePacket := make([]byte, len(nonce)+1) challengePacket[0] = packet.VERSION - copy(challengePacket[1:], nonce) + copy(challengePacket[1:], nonce[:]) _, err = conn.Write(challengePacket) if err != nil { @@ -275,7 +261,7 @@ func handleAuth(conn net.Conn, nonce []byte) (ed25519.PublicKey, error) { pubKey := ed25519.PublicKey(challengeResponsePacket[1 : 1+ed25519.PublicKeySize]) signature := ed25519.PrivateKey(challengeResponsePacket[1+ed25519.PublicKeySize:]) - if ok := ed25519.Verify(pubKey, nonce, signature); !ok { + if ok := ed25519.Verify(pubKey, nonce[:], signature); !ok { return nil, errors.New("signature verification failed") } @@ -329,8 +315,6 @@ func timeout[T packet.Payload]( case response := <-responseChan: return response case <-ctx.Done(): - sess, ok := session.FromContext(ctx) - assert.Assert(ok, "session should exist") log.Println(sess.Addr(), "timeout of", request.Type(), "request") return &packet.Error{Error: "request timeout"} } diff --git a/internal/server/session/session.go b/internal/server/session/session.go index 863b3c3..1f70ec7 100644 --- a/internal/server/session/session.go +++ b/internal/server/session/session.go @@ -22,18 +22,17 @@ type SessionManager interface { } type Session struct { - // Channel to directly write packets to the cient. - // Can be nil in cases where the connection is not available. + manager SessionManager + addr *net.TCPAddr WriteQueue chan packet.Packet - PubKey ed25519.PublicKey - manager SessionManager - addr *net.TCPAddr - id snowflake.ID - - mu sync.Mutex - challenge []byte issuedTime time.Time + challenge []byte + + PubKey ed25519.PublicKey + id snowflake.ID + + mu sync.Mutex } func NewSession(manager SessionManager, addr *net.TCPAddr, id snowflake.ID, pubKey ed25519.PublicKey) *Session { @@ -72,15 +71,11 @@ func (s *Session) Challenge() []byte { return s.challenge } -type key struct{} - -var sessKey key - -func NewContext(ctx context.Context, sess *Session) context.Context { - return context.WithValue(ctx, sessKey, sess) -} - -func FromContext(ctx context.Context) (*Session, bool) { - sess, ok := ctx.Value(sessKey).(*Session) - return sess, ok +func (s *Session) Write(ctx context.Context, pkt packet.Packet) bool { + select { + case s.WriteQueue <- pkt: + return true + case <-ctx.Done(): + return false + } } -- cgit v1.3.1