summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
authorKyren223 <Kyren223@proton.me>2024-11-21 19:57:41 +0200
committerKyren223 <Kyren223@proton.me>2024-11-21 19:57:41 +0200
commitae8e313c4829df701042e298ccb4906c43c2ce6f (patch)
tree7c679330cbad3d9fb9e8b1677518acb45de5f3f0 /internal
parent168c10aed0aab1194779481561fd73b33580bfa2 (diff)
feat: overhauled gateway
Diffstat (limited to 'internal')
-rw-r--r--internal/client/gateway/gateway.go154
1 files changed, 78 insertions, 76 deletions
diff --git a/internal/client/gateway/gateway.go b/internal/client/gateway/gateway.go
index 56f1ade..7e9ab82 100644
--- a/internal/client/gateway/gateway.go
+++ b/internal/client/gateway/gateway.go
@@ -6,11 +6,8 @@ import (
"crypto/tls"
"crypto/x509"
_ "embed"
- "errors"
- "io"
"log"
"net"
- "os"
"sync"
"time"
@@ -29,8 +26,9 @@ var (
asyncResponses []chan packet.Payload
responsesMu sync.Mutex
- connection net.Conn
- connMu sync.Mutex
+ framer packet.PacketFramer
+ conn net.Conn
+ writeMu sync.Mutex
)
func init() {
@@ -45,41 +43,68 @@ func init() {
}
}
-func Connect(ctx context.Context, program *tea.Program, privKey ed25519.PrivateKey) {
- conn, err := tls.Dial("tcp4", ":7223", tlsConfig)
- if err != nil {
- assert.NoError(err, "TODO handle error")
- }
- log.Println("established connection with server")
+func Connect(ctx context.Context, program *tea.Program, privKey ed25519.PrivateKey) error {
+ assert.Assert(conn == nil, "cannot connect, connection is active")
+
+ connChan := make(chan net.Conn, 1)
+ errChan := make(chan error, 1)
+ go func() {
+ framer = packet.NewFramer()
+ connection, err := tls.Dial("tcp4", ":7223", tlsConfig)
+ if err != nil {
+ errChan <- err
+ return
+ }
+ log.Println("established connection with server")
+
+ if err := handleAuth(ctx, privKey); err != nil {
+ errChan <- err
+ return
+ }
+ log.Println("successfully authenticated with server")
+ connChan <- connection
+ }()
- if err := handleAuth(ctx, conn, privKey); err != nil {
- assert.NoError(err, "TODO handle error")
+ select {
+ case connection := <-connChan:
+ conn = connection
+ case err := <-errChan:
+ return err
+ case <-ctx.Done():
+ return ctx.Err()
}
- log.Println("successfully authenticated with server")
- framer := packet.NewFramer()
+ go readUntilDisconnected()
+ go handlePacketStream(program)
- go func() {
- connection = conn
- handleConnection(ctx, conn, framer)
- close(framer.Out)
- conn.Close()
- connection = nil
- }()
+ return nil
+}
- go handlePacketStream(framer, program)
+func Disconnect() {
+ assert.Assert(conn != nil, "cannot disconnect, connection is inactive")
+ close(framer.Out)
+ conn.Close()
+ conn = nil
+ responsesMu.Lock()
+ for _, responseChan := range asyncResponses {
+ close(responseChan)
+ }
+ asyncResponses = nil
+ responsesMu.Unlock()
}
-func handleAuth(ctx context.Context, conn net.Conn, privKey ed25519.PrivateKey) error {
+func handleAuth(ctx context.Context, privKey ed25519.PrivateKey) error {
const nonceSize = 32
challengeRequest := make([]byte, 1+nonceSize)
- err := conn.SetDeadline(time.Now().Add(10 * time.Second))
+ deadline, _ := ctx.Deadline()
+ err := conn.SetDeadline(deadline)
assert.NoError(err, "setting deadline should not error")
defer func() {
err := conn.SetDeadline(time.Time{})
assert.NoError(err, "unsetting deadline should not error")
}()
+
bytesRead := 0
for bytesRead < 1+nonceSize {
n, err := conn.Read(challengeRequest[bytesRead:])
@@ -106,35 +131,41 @@ func handleAuth(ctx context.Context, conn net.Conn, privKey ed25519.PrivateKey)
return nil
}
-func handleConnection(ctx context.Context, conn net.Conn, framer packet.PacketFramer) {
- buffer := make([]byte, 512)
- for {
- err := conn.SetReadDeadline(time.Now().Add(time.Second))
- assert.NoError(err, "setting a read deadline should not error")
- n, err := conn.Read(buffer)
- deadlineExceeded := errors.Is(err, os.ErrDeadlineExceeded)
- if err != nil && !deadlineExceeded {
- if !errors.Is(err, io.EOF) {
- log.Println("read error:", err)
+func readUntilDisconnected() {
+ ctx, cancel := context.WithCancel(context.Background())
+ defer cancel()
+ go func() {
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ default:
+ if conn == nil {
+ cancel()
+ return
+ }
}
- break
}
+ }()
- if ctx.Err() != nil {
- log.Println("context error:", ctx.Err())
+ buffer := make([]byte, 512)
+ for conn != nil {
+ n, err := conn.Read(buffer)
+ if err != nil {
+ log.Println("server connectivity error: ", err)
break
}
err = framer.Push(ctx, buffer[:n])
if ctx.Err() != nil {
- log.Println("context error:", ctx.Err())
+ log.Println("server connectivity error: ", ctx.Err())
break
}
- assert.NoError(err, "packets from server should always be correct")
+ assert.NoError(err, "packets from server should always be correctly formatted")
}
}
-func handlePacketStream(framer packet.PacketFramer, program *tea.Program) {
+func handlePacketStream(program *tea.Program) {
for {
pkt, ok := <-framer.Out
if !ok {
@@ -163,58 +194,29 @@ func handlePacketStream(framer packet.PacketFramer, program *tea.Program) {
}
}
-func conn() net.Conn {
- return connection
-}
-
func Send(request packet.Payload) <-chan packet.Payload {
responseChan := make(chan packet.Payload)
go func() {
pkt := packet.NewPacket(packet.NewMsgPackEncoder(request))
- conn := conn()
+ responsesMu.Lock()
+ asyncResponses = append(asyncResponses, responseChan)
+ responsesMu.Unlock()
+
if conn == nil {
log.Println("request send error:", "connection is closed")
close(responseChan)
return
}
- connMu.Lock()
- errDeadline := conn.SetWriteDeadline(time.Now().Add(5 * time.Second))
- assert.NoError(errDeadline, "setting a write deadline should not error")
+ writeMu.Lock()
_, err := pkt.Into(conn)
- errDeadline = conn.SetWriteDeadline(time.Time{})
- assert.NoError(errDeadline, "setting a write deadline should not error")
- connMu.Unlock()
+ writeMu.Unlock()
if err != nil {
log.Println("request send error:", err)
close(responseChan)
return
}
-
- responsesMu.Lock()
- asyncResponses = append(asyncResponses, responseChan)
- responsesMu.Unlock()
-
- // TODO: this is a really bad implementation
- // Should refactor to make it better
- // Maybe have a timeout function similar to the one on the server?
- // In any case this is like super bad to just arbitrary hang for 5 seconds
- // And then try to close the thing so it times out
- time.Sleep(5 * time.Second)
- responsesMu.Lock()
- index := -1
- for i, ch := range asyncResponses {
- if ch == responseChan {
- index = i
- }
- }
- if index != -1 {
- copy(asyncResponses[index:], asyncResponses[index+1:])
- asyncResponses = asyncResponses[:len(asyncResponses)-1]
- close(responseChan)
- }
- responsesMu.Unlock()
}()
return responseChan
}