summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/client/gateway/gateway.go2
-rw-r--r--internal/packet/packet.go25
-rw-r--r--internal/packet/packet_test.go115
-rw-r--r--internal/server/server.go8
4 files changed, 140 insertions, 10 deletions
diff --git a/internal/client/gateway/gateway.go b/internal/client/gateway/gateway.go
index 2798950..c196b1b 100644
--- a/internal/client/gateway/gateway.go
+++ b/internal/client/gateway/gateway.go
@@ -57,7 +57,7 @@ func Connect(ctx context.Context, program *tea.Program, privKey ed25519.PrivateK
}
log.Println("successfully authenticated with server")
- framer := packet.NewFramer(ctx)
+ framer := packet.NewFramer()
go func() {
connection = conn
diff --git a/internal/packet/packet.go b/internal/packet/packet.go
index 9d84486..c7986cb 100644
--- a/internal/packet/packet.go
+++ b/internal/packet/packet.go
@@ -52,6 +52,8 @@ type PacketType uint8
const (
PacketError PacketType = iota
PacketSendMessage
+ PacketPushedMessages
+ PacketGetMessageRange
PacketMessages
)
@@ -61,6 +63,10 @@ func (t PacketType) String() string {
return "PacketError"
case PacketSendMessage:
return "PacketSendMessage"
+ case PacketPushedMessages:
+ return "PacketPushedMessages"
+ case PacketGetMessageRange:
+ return "PacketGetMessageRange"
case PacketMessages:
return "PacketMessages"
default:
@@ -70,7 +76,7 @@ func (t PacketType) String() string {
func (e PacketType) IsSupported() bool {
switch e {
- case PacketError, PacketSendMessage, PacketMessages:
+ case PacketError, PacketSendMessage, PacketPushedMessages, PacketGetMessageRange, PacketMessages:
return true
default:
return false
@@ -80,10 +86,13 @@ func (e PacketType) IsSupported() bool {
// True for all packets that a server may push passively to the client.
func (e PacketType) IsPush() bool {
switch e {
- case PacketError, PacketSendMessage:
+ case PacketError, PacketSendMessage, PacketGetMessageRange, PacketMessages:
return false
- default:
+ case PacketPushedMessages:
return true
+ default:
+ assert.Never("should never happen")
+ return false
}
}
@@ -190,10 +199,14 @@ func (p Packet) DecodedPayload() (Payload, error) {
switch p.Type() {
case PacketError:
payload = &ErrorMessage{}
- case PacketMessages:
- payload = &Messages{}
+ case PacketPushedMessages:
+ payload = &PushedMessages{}
case PacketSendMessage:
payload = &SendMessage{}
+ case PacketGetMessageRange:
+ payload = &GetMessagesRange{}
+ case PacketMessages:
+ payload = &Messages{}
default:
assert.Never("packet type of a packet struct must always be valid")
}
@@ -212,7 +225,7 @@ type PacketFramer struct {
Out chan Packet
}
-func NewFramer(ctx context.Context) PacketFramer {
+func NewFramer() PacketFramer {
return PacketFramer{
Out: make(chan Packet, 10),
}
diff --git a/internal/packet/packet_test.go b/internal/packet/packet_test.go
new file mode 100644
index 0000000..e5a3612
--- /dev/null
+++ b/internal/packet/packet_test.go
@@ -0,0 +1,115 @@
+package packet
+
+import (
+ "context"
+ "reflect"
+ "slices"
+ "testing"
+ "time"
+
+ "github.com/stretchr/testify/require"
+
+ "github.com/kyren223/eko/internal/data"
+ "github.com/kyren223/eko/pkg/snowflake"
+)
+
+func TestPacketEncodingDecoding(t *testing.T) {
+ testPacketEncodingDecoding(t, &ErrorMessage{"Hello, World!"})
+
+ node := snowflake.NewNode(1)
+ id := node.Generate()
+ testPacketEncodingDecoding(t, &Messages{Messages: []data.Message{
+ {
+ ID: node.Generate(),
+ SenderID: node.Generate(),
+ Content: "MyMessage",
+ FrequencyID: &id,
+ ReceiverID: nil,
+ },
+ {
+ ID: node.Generate(),
+ SenderID: node.Generate(),
+ Content: "Another Message\nWith a bunch of stuff",
+ FrequencyID: nil,
+ ReceiverID: &id,
+ },
+ }})
+}
+
+func testPacketEncodingDecoding(t *testing.T, payload Payload) {
+ t.Helper()
+
+ jsonEncoder := NewJsonEncoder(payload)
+ jsonPacket := NewPacket(jsonEncoder)
+
+ msgPackEncoder := NewJsonEncoder(payload)
+ msgPackPacket := NewPacket(msgPackEncoder)
+
+ jsonPayload, err := jsonPacket.DecodedPayload()
+ require.NoError(t, err, "json payload decoding should not fail")
+ require.True(t, reflect.DeepEqual(payload, jsonPayload), "json payload mismatch", "got", jsonPayload, "want", payload)
+
+ msgPackPayload, err := msgPackPacket.DecodedPayload()
+ require.NoError(t, err, "msgpack payload decoding should not fail")
+ require.True(t, reflect.DeepEqual(payload, msgPackPayload), "msgpack payload mismatch", "got", msgPackPayload, "want", payload)
+}
+
+func TestPacketFramer(t *testing.T) {
+ ctx, cancel := context.WithTimeout(context.Background(), time.Millisecond*20)
+ defer cancel()
+ framer := NewFramer()
+
+ pkt := NewPacket(NewJsonEncoder(&ErrorMessage{"Hello, World!"}))
+ length := len(pkt.data)
+ count := 5
+ data := make([]byte, length*count)
+ for i := 0; i < count; i++ {
+ copy(data[i*length:], pkt.data)
+ }
+
+ first := 2
+ require.True(t, first < HEADER_SIZE, "TEST ERROR size needs to be less than the header")
+ err := framer.Push(ctx, data[:first])
+ data = data[first:]
+ require.NoError(t, err, "expecting framer to return nil and wait for more data for header")
+ require.False(t, doesChannelHaveValue(framer.Out), "expecting channel to block with no value")
+
+ err = framer.Push(ctx, data[:HEADER_SIZE])
+ data = data[HEADER_SIZE:]
+ require.NoError(t, err, "expecting framer to return nil and wait for more data for payload")
+ require.False(t, doesChannelHaveValue(framer.Out), "expecting channel to block with no value")
+
+ err = framer.Push(ctx, data[:length])
+ data = data[length:]
+ require.NoError(t, err, "expecting framer to return nil and process exactly 1 packet")
+ select {
+ case p, ok := <-framer.Out:
+ require.True(t, ok, "expecting a value from the framer")
+ require.True(t, slices.Equal(p.data, pkt.data), "expecting packet to be equal", "got", p.data, "want", pkt.data)
+ default:
+ require.Fail(t, "expected packet but channel was blocking")
+ }
+ require.False(t, doesChannelHaveValue(framer.Out), "expecting channel to block with no value because it was consumed already")
+
+ err = framer.Push(ctx, data[:])
+ require.NoError(t, err, "expecting framer to return nil and process the rest of the packets")
+ for i := 0; i < count-1; i++ {
+ select {
+ case p, ok := <-framer.Out:
+ require.True(t, ok, "expecting a value from the framer")
+ require.True(t, slices.Equal(p.data, pkt.data), "expecting packet to be equal", "got", p.data, "want", pkt.data)
+ default:
+ require.Failf(t, "channel blocked", "%v expected packet", i)
+ }
+ }
+ require.False(t, doesChannelHaveValue(framer.Out), "expecting channel to block with no value because it was consumed already")
+}
+
+func doesChannelHaveValue[T any](c <-chan T) bool {
+ select {
+ case _, ok := <-c:
+ return ok
+ default:
+ return false
+ }
+}
diff --git a/internal/server/server.go b/internal/server/server.go
index 05af434..1b8dccb 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -152,7 +152,7 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
sess := session.NewSession(server, addr, id, pubKey)
server.AddSession(sess)
ctx = session.NewContext(ctx, sess)
- framer := packet.NewFramer(ctx)
+ framer := packet.NewFramer()
defer func() {
conn.Close()
@@ -292,9 +292,11 @@ func processRequest(ctx context.Context, request packet.Payload) packet.Payload
switch request := request.(type) {
case *packet.SendMessage:
- return timeout(100 * time.Millisecond, api.SendMessage, ctx, request)
+ return timeout(50 * time.Millisecond, api.SendMessage, ctx, request)
+ case *packet.GetMessagesRange:
+ return timeout(100 * time.Millisecond, api.GetMessages, ctx, request)
default:
- return &packet.ErrorMessage{Error: "use of unsupported packet type"}
+ return &packet.ErrorMessage{Error: "use of disallowed packet type for request"}
}
}