summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
authorKyren223 <Kyren223@proton.me>2024-11-27 12:52:22 +0200
committerKyren223 <Kyren223@proton.me>2024-11-27 12:52:22 +0200
commit6a34d7be01256df472015564a655d8253acd8724 (patch)
tree0956f3d776ef7963d569e77d7c68324f715b5693 /internal
parente415a13211316cb804752c18833a10057bc0994c (diff)
Refactored server infrastructure
Diffstat (limited to 'internal')
-rw-r--r--internal/packet/packet.go23
-rw-r--r--internal/server/api/database.go6
-rw-r--r--internal/server/server.go68
-rw-r--r--internal/server/session/session.go35
4 files changed, 59 insertions, 73 deletions
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
+ }
}