diff options
Diffstat (limited to 'internal/client/gateway/gateway.go')
| -rw-r--r-- | internal/client/gateway/gateway.go | 105 |
1 files changed, 50 insertions, 55 deletions
diff --git a/internal/client/gateway/gateway.go b/internal/client/gateway/gateway.go index 9151d74..a8f4241 100644 --- a/internal/client/gateway/gateway.go +++ b/internal/client/gateway/gateway.go @@ -27,9 +27,6 @@ var certPEM []byte var ( tlsConfig *tls.Config - asyncResponses []chan packet.Payload - responsesMu sync.Mutex - framer packet.PacketFramer conn net.Conn writeMu sync.Mutex @@ -81,13 +78,13 @@ func connect(ctx context.Context, privKey ed25519.PrivateKey) (snowflake.ID, err errChan <- err return } - log.Println("established connection with server") + log.Println("established connection with the server") if id, err = handleAuth(ctx, connection, privKey); err != nil { errChan <- err return } - log.Println("successfully authenticated with server") + log.Println("successfully authenticated with the server") connChan <- connection }() @@ -100,7 +97,7 @@ func connect(ctx context.Context, privKey ed25519.PrivateKey) (snowflake.ID, err return 0, ctx.Err() } - go readForever() + go readForever(conn) go handlePacketStream() return id, nil @@ -108,7 +105,8 @@ func connect(ctx context.Context, privKey ed25519.PrivateKey) (snowflake.ID, err func handleAuth(ctx context.Context, conn net.Conn, privKey ed25519.PrivateKey) (snowflake.ID, error) { const nonceSize = 32 - challengeRequest := make([]byte, 1+nonceSize) + const packetSize = 1 + nonceSize // For version byte + challengeRequest := make([]byte, packetSize) deadline, _ := ctx.Deadline() err := conn.SetDeadline(deadline) @@ -119,7 +117,7 @@ func handleAuth(ctx context.Context, conn net.Conn, privKey ed25519.PrivateKey) }() bytesRead := 0 - for bytesRead < 1+nonceSize { + for bytesRead < packetSize { n, err := conn.Read(challengeRequest[bytesRead:]) if err != nil { return 0, err @@ -155,9 +153,9 @@ func handleAuth(ctx context.Context, conn net.Conn, privKey ed25519.PrivateKey) return id, nil } -func readForever() { +func readForever(conn net.Conn) { buffer := make([]byte, 512) - for conn != nil { + for { n, err := conn.Read(buffer) if err != nil { onDisconnect(err) @@ -168,7 +166,7 @@ func readForever() { err = framer.Push(ctx, buffer[:n]) if ctx.Err() != nil { cancel() - onDisconnect(errors.New("framer blocked for more than a second, closing connection")) + onDisconnect(errors.New("new packet blocked for more than a second, closing connection")) return } cancel() @@ -186,41 +184,27 @@ func handlePacketStream() { payload, err := pkt.DecodedPayload() assert.NoError(err, "server should always provide a decodeable packet") - if pkt.Type().IsPush() { - log.Println("received streamed packet:", payload) - ui.Program.Send(payload) - continue - } - - responsesMu.Lock() - assert.Assert(len(asyncResponses) != 0, "there must always be at least 1 response waiting") - responseChan := asyncResponses[0] - copy(asyncResponses, asyncResponses[1:]) - asyncResponses = asyncResponses[:len(asyncResponses)-1] - responsesMu.Unlock() - - go func() { - responseChan <- payload - }() + log.Println("received streamed packet:", payload) + ui.Program.Send(payload) } } func Disconnect() { - assert.Assert(conn != nil, "cannot disconnect, connection is inactive") - conn.Close() - closed = true + if conn != nil { + conn.Close() + closed = true + } } func onDisconnect(err error) { + writeMu.Lock() + defer writeMu.Unlock() + if conn == nil { + return + } conn.Close() conn = nil close(framer.Out) - responsesMu.Lock() - for _, responseChan := range asyncResponses { - close(responseChan) - } - asyncResponses = nil - responsesMu.Unlock() if closed { log.Println("connection closed") ui.Program.Send(ConnectionClosed{}) @@ -230,27 +214,38 @@ func onDisconnect(err error) { } } -func Send(request packet.Payload) <-chan packet.Payload { - responseChan := make(chan packet.Payload, 1) - go func() { - pkt := packet.NewPacket(packet.NewMsgPackEncoder(request)) - - responsesMu.Lock() - asyncResponses = append(asyncResponses, responseChan) - responsesMu.Unlock() - - if conn == nil { - log.Println("request send error:", "connection is closed") - return - } +type RequestSentMsg struct { + request packet.Payload + err error +} - writeMu.Lock() - _, err := pkt.Into(conn) - writeMu.Unlock() +func Send(request packet.Payload) tea.Cmd { + return func() tea.Msg { + err := send(request) if err != nil { log.Println("request send error:", err) - return } - }() - return responseChan + return RequestSentMsg{ + request: request, + err: err, + } + } +} + +func send(request packet.Payload) error { + pkt := packet.NewPacket(packet.NewMsgPackEncoder(request)) + + writeMu.Lock() + if conn == nil { + writeMu.Unlock() + return errors.New("connection is closed") + } + _, err := pkt.Into(conn) + writeMu.Unlock() + + if err != nil { + return err + } + + return nil } |
