mirror of
https://github.com/Kyren223/eko.git
synced 2026-08-30 19:31:30 +00:00
test(packet): test packet framer and packet encoding/decoding
This commit is contained in:
5
go.mod
5
go.mod
@@ -6,6 +6,7 @@ require (
|
||||
github.com/charmbracelet/bubbles v0.20.0
|
||||
github.com/charmbracelet/bubbletea v1.1.1
|
||||
github.com/charmbracelet/lipgloss v0.13.0
|
||||
github.com/stretchr/testify v1.7.0
|
||||
github.com/vmihailenco/msgpack/v5 v5.4.1
|
||||
)
|
||||
|
||||
@@ -14,6 +15,7 @@ require (
|
||||
github.com/aymanbagabas/go-osc52/v2 v2.0.1 // indirect
|
||||
github.com/charmbracelet/x/ansi v0.2.3 // indirect
|
||||
github.com/charmbracelet/x/term v0.2.0 // indirect
|
||||
github.com/davecgh/go-spew v1.1.0 // indirect
|
||||
github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f // indirect
|
||||
github.com/lucasb-eyer/go-colorful v1.2.0 // indirect
|
||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||
@@ -22,10 +24,11 @@ require (
|
||||
github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6 // indirect
|
||||
github.com/muesli/cancelreader v0.2.2 // indirect
|
||||
github.com/muesli/termenv v0.15.2 // indirect
|
||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||
github.com/rivo/uniseg v0.4.7 // indirect
|
||||
github.com/stretchr/testify v1.7.0 // indirect
|
||||
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
|
||||
golang.org/x/sync v0.8.0 // indirect
|
||||
golang.org/x/sys v0.24.0 // indirect
|
||||
golang.org/x/text v0.14.0 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c // indirect
|
||||
)
|
||||
|
||||
1
go.sum
1
go.sum
@@ -52,6 +52,7 @@ golang.org/x/sys v0.24.0 h1:Twjiwq9dn6R1fQcyiK+wQyHWfaz/BJB+YIpzU/Cv3Xg=
|
||||
golang.org/x/sys v0.24.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/text v0.14.0 h1:ScX5w1eTa3QqT8oi6+ziP7dTV1S2+ALU0bI+0zXKWiQ=
|
||||
golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c h1:dUUwHk2QECo/6vqA44rthZ8ie2QXMNeKRTHCNY2nXvo=
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
|
||||
115
internal/packet/packet_test.go
Normal file
115
internal/packet/packet_test.go
Normal file
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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"}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user