mirror of
https://github.com/Kyren223/eko.git
synced 2026-07-21 02:51:07 +00:00
Fixed server hanging until the client closes the connection,
it now forcibly closes the connection when it's terminated
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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[:])
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user