summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--internal/server/api/api.go5
-rw-r--r--internal/server/server.go75
-rw-r--r--internal/server/session/session.go15
3 files changed, 58 insertions, 37 deletions
diff --git a/internal/server/api/api.go b/internal/server/api/api.go
index 7dcf1ab..fba6cdf 100644
--- a/internal/server/api/api.go
+++ b/internal/server/api/api.go
@@ -77,10 +77,7 @@ func CreateOrGetUser(ctx context.Context, node *snowflake.Node, pubKey ed25519.P
PublicKey: pubKey,
})
}
- if err != nil {
- return data.User{}, err
- }
- return user, nil
+ return user, err
}
func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.CreateNetwork) packet.Payload {
diff --git a/internal/server/server.go b/internal/server/server.go
index c15414f..db3cb6b 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -46,6 +46,7 @@ func init() {
}
type server struct {
+ ctx context.Context
node *snowflake.Node
sessions map[snowflake.ID]*session.Session
sessMu sync.RWMutex
@@ -54,15 +55,16 @@ type server struct {
// Creates a new server on the given port.
// Will generate a unique node ID automatically, will crash if there are no available IDs.
-func NewServer(port uint16) server {
+func NewServer(ctx context.Context, port uint16) server {
assert.Assert(nodeId <= snowflake.NodeMax, "maximum amount of servers reached")
node := snowflake.NewNode(nodeId)
nodeId++
return server{
+ ctx: ctx,
node: node,
- Port: port,
sessions: map[snowflake.ID]*session.Session{},
+ Port: port,
}
}
@@ -91,7 +93,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(ctx context.Context) error {
+func (s *server) Run() error {
listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(int(s.Port)), tlsConfig)
if err != nil {
log.Fatalf("error starting server: %s", err)
@@ -100,7 +102,7 @@ func (s *server) Run(ctx context.Context) error {
assert.AddFlush(listener)
defer listener.Close()
go func() {
- <-ctx.Done()
+ <-s.ctx.Done()
listener.Close()
}()
@@ -116,7 +118,7 @@ func (s *server) Run(ctx context.Context) error {
}
wg.Add(1)
go func() {
- s.handleConnection(ctx, conn)
+ s.handleConnection(conn)
wg.Done()
}()
}
@@ -128,22 +130,32 @@ func (s *server) Run(ctx context.Context) error {
return nil
}
-func (server *server) handleConnection(ctx context.Context, conn net.Conn) {
+func (server *server) handleConnection(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")
+ initialCtx, cancel := context.WithTimeout(server.ctx, 5*time.Second)
+ deadline, _ := initialCtx.Deadline()
+ err := conn.SetDeadline(deadline)
+ assert.NoError(err, "setting read deadline should not error")
+
+ err = conn.SetDeadline(time.Time{})
+ assert.NoError(err, "unsetting read deadline should not error")
+
pubKey, err := handleAuth(conn)
if err != nil {
+ cancel()
log.Println(addr, err)
conn.Close()
log.Println(addr, "disconnected")
return
}
- user, err := api.CreateOrGetUser(ctx, server.Node(), pubKey)
+ user, err := api.CreateOrGetUser(initialCtx, server.Node(), pubKey)
if err != nil {
+ cancel()
log.Println(addr, "user creation/fetching error:", err)
conn.Close()
log.Println(addr, "disconnected")
@@ -158,42 +170,53 @@ func (server *server) handleConnection(ctx context.Context, conn net.Conn) {
binary.BigEndian.PutUint64(id[:], uint64(user.ID))
_, err = conn.Write(id[:])
if err != nil {
+ cancel()
log.Println(addr, "failed to write user id")
conn.Close()
log.Println(addr, "disconnected")
return
}
+ cancel()
+
+ go func() {
+ <-server.ctx.Done()
+ conn.Close()
+ }()
defer func() {
conn.Close()
server.RemoveSession(sess.ID())
- close(framer.Out) // avoid leaking goroutine
log.Println(addr, "disconnected")
}()
go func() {
for {
- packet, ok := <-sess.WriteQueue
+ packet, ok := sess.Read(server.ctx)
if !ok {
- break
+ return
}
log.Println(addr, "sending packet:", packet)
if _, err := packet.Into(conn); err != nil {
log.Println(addr, err)
- break
+ return
}
}
}()
go func() {
for {
- request, ok := <-framer.Out
- if !ok {
- break
- }
- response := processPacket(ctx, sess, request)
- if ok := sess.Write(ctx, response); !ok {
- break
+ select {
+ case <-server.ctx.Done():
+ return
+ case request, ok := <-framer.Out:
+ if !ok {
+ return
+ }
+
+ response := processPacket(server.ctx, sess, request)
+ if ok := sess.Write(server.ctx, response); !ok {
+ return
+ }
}
}
}()
@@ -208,15 +231,15 @@ func (server *server) handleConnection(ctx context.Context, conn net.Conn) {
break
}
- err = framer.Push(ctx, buffer[:n])
- if ctx.Err() != nil {
- log.Println(addr, ctx.Err())
+ err = framer.Push(server.ctx, buffer[:n])
+ if server.ctx.Err() != nil {
+ log.Println(addr, server.ctx.Err())
break
}
if err != nil {
payload := packet.Error{Error: err.Error()}
pkt := packet.NewPacket(packet.NewMsgPackEncoder(&payload))
- sess.Write(ctx, pkt)
+ sess.Write(server.ctx, pkt)
break
}
}
@@ -227,14 +250,6 @@ func handleAuth(conn net.Conn) (ed25519.PublicKey, error) {
_, 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() {
- err := conn.SetDeadline(time.Time{})
- assert.NoError(err, "unsetting read deadline should not error")
- }()
-
challengePacket := make([]byte, len(nonce)+1)
challengePacket[0] = packet.VERSION
copy(challengePacket[1:], nonce[:])
diff --git a/internal/server/session/session.go b/internal/server/session/session.go
index 1f70ec7..dea9af2 100644
--- a/internal/server/session/session.go
+++ b/internal/server/session/session.go
@@ -24,7 +24,7 @@ type SessionManager interface {
type Session struct {
manager SessionManager
addr *net.TCPAddr
- WriteQueue chan packet.Packet
+ writeQueue chan packet.Packet
issuedTime time.Time
challenge []byte
@@ -37,7 +37,7 @@ type Session struct {
func NewSession(manager SessionManager, addr *net.TCPAddr, id snowflake.ID, pubKey ed25519.PublicKey) *Session {
session := &Session{
- WriteQueue: make(chan packet.Packet, 10),
+ writeQueue: make(chan packet.Packet, 10),
PubKey: pubKey,
manager: manager,
addr: addr,
@@ -73,9 +73,18 @@ func (s *Session) Challenge() []byte {
func (s *Session) Write(ctx context.Context, pkt packet.Packet) bool {
select {
- case s.WriteQueue <- pkt:
+ case s.writeQueue <- pkt:
return true
case <-ctx.Done():
return false
}
}
+
+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
+ }
+}