summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/client/api/api.go29
-rw-r--r--internal/client/client.go226
-rw-r--r--internal/client/gateway/gateway.go215
-rw-r--r--internal/client/gateway/server.crt (renamed from internal/client/server.crt)0
-rw-r--r--internal/client/ui.go125
-rw-r--r--internal/data/chat.go29
-rw-r--r--internal/packet/encoders.go17
-rw-r--r--internal/packet/packet.go37
-rw-r--r--internal/server/server.go74
9 files changed, 495 insertions, 257 deletions
diff --git a/internal/client/api/api.go b/internal/client/api/api.go
new file mode 100644
index 0000000..56cf054
--- /dev/null
+++ b/internal/client/api/api.go
@@ -0,0 +1,29 @@
+package api
+
+import (
+ "log"
+
+ tea "github.com/charmbracelet/bubbletea"
+
+ "github.com/kyren223/eko/internal/client/gateway"
+ "github.com/kyren223/eko/internal/packet"
+ "github.com/kyren223/eko/pkg/assert"
+)
+
+func SendMessage(message string) tea.Cmd {
+ return func() tea.Msg {
+ log.Println("request SendMessage sent")
+ request := packet.SendMessage{Content: message}
+ responsePayload, ok := <-gateway.Send(&request)
+ if !ok {
+ log.Println("request timeout")
+ return ""
+ }
+ log.Println("request SendMessage received response")
+
+ response := responsePayload.(*packet.ErrorMessage)
+ assert.Assert(ok, "server should return an ErrorMessage response for a SendMessage request")
+ assert.Assert(response.IsOk(), "server should always return an OK if we didn't mess up")
+ return ""
+ }
+}
diff --git a/internal/client/client.go b/internal/client/client.go
index 5f42c8f..1276f01 100644
--- a/internal/client/client.go
+++ b/internal/client/client.go
@@ -2,125 +2,167 @@ package client
import (
"context"
- "crypto/tls"
- "crypto/x509"
- _ "embed"
+ "crypto/ed25519"
"fmt"
"log"
- "os"
- "time"
+ "slices"
+ "strings"
+ "github.com/charmbracelet/bubbles/cursor"
+ "github.com/charmbracelet/bubbles/textarea"
+ "github.com/charmbracelet/bubbles/viewport"
+ tea "github.com/charmbracelet/bubbletea"
+ "github.com/charmbracelet/lipgloss"
+
+ "github.com/kyren223/eko/internal/client/api"
+ "github.com/kyren223/eko/internal/client/gateway"
+ "github.com/kyren223/eko/internal/data"
"github.com/kyren223/eko/internal/packet"
"github.com/kyren223/eko/pkg/assert"
- "github.com/vmihailenco/msgpack/v5"
)
-//go:embed server.crt
-var certPEM []byte
+type BubbleTeaCloser struct {
+ program *tea.Program
+}
-var tlsConfig *tls.Config
+func (c BubbleTeaCloser) Close() error {
+ c.program.Quit()
+ return nil
+}
func Run() {
- logFile, err := os.OpenFile("client.log", os.O_APPEND|os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0666)
- if err != nil {
- log.Fatal(err)
- }
- defer logFile.Close()
- log.SetOutput(logFile)
+ log.Println("client started")
+ program := tea.NewProgram(initialModel(), tea.WithAltScreen())
+ assert.AddFlush(BubbleTeaCloser{program})
- certPool := x509.NewCertPool()
- if !certPool.AppendCertsFromPEM(certPEM) {
- log.Fatalln("failed to append server certificate")
- }
+ _, privKey, err := ed25519.GenerateKey(nil)
+ assert.NoError(err, "private key gen should not error")
- tlsConfig = &tls.Config{
- RootCAs: certPool,
- ServerName: "localhost",
+ gateway.Connect(context.Background(), program, privKey)
+ if _, err := program.Run(); err != nil {
+ log.Println(err)
}
+}
- log.Println("client started, waiting for user input...")
- startUI()
-
- // for {
- // fmt.Print("> ")
- // input, _ := bufio.NewReader(os.Stdin).ReadString('\n')
- // input = strings.TrimSpace(input)
- // if input == ":q" || input == "exit" || input == "quit" {
- // break
- // }
- // if input == "" {
- // continue
- // }
- // err := processRequest(input)
- // if err != nil {
- // log.Println(err)
- // fmt.Println(err)
- // }
- // }
+type model struct {
+ viewport viewport.Model
+ messages []string
+ textarea textarea.Model
+ senderStyle lipgloss.Style
+ err error
}
-func sendMessage(message string) error {
- request := packet.SendMessageMessage{Content: message}
- var response packet.EkoMessage
- if err := SendAndReceive(&request, &response); err != nil {
- return err
+func initialModel() model {
+ ta := textarea.New()
+ ta.Placeholder = "Send a message..."
+ ta.Focus()
+
+ ta.Prompt = "┃ "
+ ta.CharLimit = 280
+
+ ta.SetWidth(30)
+ ta.SetHeight(3)
+
+ // Remove cursor line styling
+ ta.FocusedStyle.CursorLine = lipgloss.NewStyle()
+
+ ta.ShowLineNumbers = false
+
+ vp := viewport.New(30, 20)
+ // vp.SetContent("Welcome to Eko!\n Type a message and press Enter to send.")
+
+ ta.KeyMap.InsertNewline.SetEnabled(false)
+
+ return model{
+ textarea: ta,
+ messages: []string{},
+ viewport: vp,
+ senderStyle: lipgloss.NewStyle().Foreground(lipgloss.Color("5")),
+ err: nil,
}
- assert.Assert(response.Message == "Eko OK", "server should return an OK status")
- return nil
}
-func getMessages() ([]string, error) {
- request := packet.GetMessagesMessage{}
- var response packet.MessagesMessage
- if err := SendAndReceive(&request, &response); err != nil {
- return nil, err
- }
- var messages []string
- for _, message := range response.Messages {
- messages = append(messages, message.Contents)
- }
- return messages, nil
+func (m model) Init() tea.Cmd {
+ return textarea.Blink
}
-func SendAndReceive(request packet.TypedMessage, response packet.TypedMessage) error {
- conn, err := tls.Dial("tcp4", ":7223", tlsConfig)
- if err != nil {
- return fmt.Errorf("error establishing connection with server: %v", err)
- }
- defer conn.Close()
- log.Println("established connection with server:", conn.RemoteAddr().String())
+func (m model) View() string {
+ return fmt.Sprintf(
+ "%s\n%s",
+ m.viewport.View(),
+ m.textarea.View(),
+ ) + ""
+}
- encoder, err := packet.NewMsgPackEncoder(request)
- if err != nil {
- return fmt.Errorf("error encoding request: %v", err)
- }
- requestPacket := packet.NewPacket(encoder)
- err = requestPacket.Into(conn)
- if err != nil {
- return fmt.Errorf("error sending request: %v", err)
- }
- log.Println("sent request to server")
+func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
+ switch msg := msg.(type) {
+ case tea.WindowSizeMsg:
+ m.viewport.Width = msg.Width
+ m.viewport.Height = msg.Height - m.textarea.Height()
+ m.viewport.GotoBottom()
+
+ m.textarea.SetWidth(msg.Width)
+ log.Println("resized to:", msg.Width, "x", msg.Height)
- ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
- defer cancel()
- out, outErr := packet.RunFramer(ctx, conn)
+ var vpCmd tea.Cmd
+ var taCmd tea.Cmd
+ m.viewport, vpCmd = m.viewport.Update(msg)
+ m.textarea, taCmd = m.textarea.Update(msg)
+ return m, tea.Batch(vpCmd, taCmd)
- select {
- case responsePacket := <-out:
- if err := responsePacket.DecodePayload(response); err != nil {
- if responsePacket.Type() != packet.PacketError {
- return fmt.Errorf("error decoding response: %v", err)
+ case tea.KeyMsg:
+ switch msg.Type {
+ case tea.KeyEsc, tea.KeyCtrlC:
+ return m, tea.Quit
+
+ case tea.KeyEnter:
+ value := m.textarea.Value()
+
+ content := strings.TrimSpace(value)
+ if content == "" {
+ return m, nil
}
- var errorResponse packet.ErrorMessage
- if err := responsePacket.DecodePayload(&errorResponse); err != nil {
- return fmt.Errorf("error decoding error packet: %w", err)
+
+ m.textarea.Reset()
+ return m, api.SendMessage(content)
+
+ default:
+ // Send all other keypresses to the textarea.
+ var cmd tea.Cmd
+ m.textarea, cmd = m.textarea.Update(msg)
+ return m, cmd
+ }
+
+ case packet.Payload:
+ switch msg := msg.(type) {
+ case *packet.Messages:
+ slices.SortFunc(msg.Messages, func(a, b data.Message) int {
+ return a.CmpTimestamp(b)
+ })
+
+ m.messages = []string{}
+ for _, message := range msg.Messages {
+ m.messages = append(m.messages, message.Contents)
}
- return fmt.Errorf("server error: %v", errorResponse.Error)
+
+ m.viewport.SetContent(strings.Join(m.messages, "\n"))
+ m.viewport.GotoBottom()
+
+ var cmd tea.Cmd
+ m.viewport, cmd = m.viewport.Update(msg)
+ return m, cmd
+
+ default:
+ return m, nil
}
- case err := <-outErr:
- return fmt.Errorf("error receiving response: %v", err)
- }
+ case cursor.BlinkMsg:
+ // Textarea should also process cursor blinks.
+ var cmd tea.Cmd
+ m.textarea, cmd = m.textarea.Update(msg)
+ return m, cmd
- return nil
+ default:
+ return m, nil
+ }
}
diff --git a/internal/client/gateway/gateway.go b/internal/client/gateway/gateway.go
new file mode 100644
index 0000000..2798950
--- /dev/null
+++ b/internal/client/gateway/gateway.go
@@ -0,0 +1,215 @@
+package gateway
+
+import (
+ "context"
+ "crypto/ed25519"
+ "crypto/tls"
+ "crypto/x509"
+ _ "embed"
+ "errors"
+ "io"
+ "log"
+ "net"
+ "os"
+ "sync"
+ "time"
+
+ tea "github.com/charmbracelet/bubbletea"
+
+ "github.com/kyren223/eko/internal/packet"
+ "github.com/kyren223/eko/pkg/assert"
+)
+
+//go:embed server.crt
+var certPEM []byte
+
+var (
+ tlsConfig *tls.Config
+
+ asyncResponses []chan packet.Payload
+ responsesMu sync.Mutex
+
+ connection net.Conn
+ connMu sync.Mutex
+)
+
+func init() {
+ certPool := x509.NewCertPool()
+ if !certPool.AppendCertsFromPEM(certPEM) {
+ log.Fatalln("failed to append server certificate")
+ }
+
+ tlsConfig = &tls.Config{
+ RootCAs: certPool,
+ ServerName: "localhost",
+ }
+}
+
+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")
+
+ if err := handleAuth(ctx, conn, privKey); err != nil {
+ assert.NoError(err, "TODO handle error")
+ }
+ log.Println("successfully authenticated with server")
+
+ framer := packet.NewFramer(ctx)
+
+ go func() {
+ connection = conn
+ handleConnection(ctx, conn, framer)
+ close(framer.Out)
+ conn.Close()
+ connection = nil
+ }()
+
+ go handlePacketStream(framer, program)
+}
+
+func handleAuth(ctx context.Context, conn net.Conn, privKey ed25519.PrivateKey) error {
+ const nonceSize = 32
+ challengeRequest := make([]byte, 1+nonceSize)
+
+ err := conn.SetDeadline(time.Now().Add(10 * time.Second))
+ 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:])
+ if err != nil {
+ return err
+ }
+ bytesRead += n
+ }
+
+ assert.Assert(challengeRequest[0] == packet.VERSION, "client should always have the same version as the server")
+
+ challengeResponse := make([]byte, 1+ed25519.PublicKeySize+ed25519.SignatureSize)
+ challengeResponse[0] = packet.VERSION
+ copy(challengeResponse[1:1+ed25519.PublicKeySize], privKey.Public().(ed25519.PublicKey))
+ signedNonce := ed25519.Sign(privKey, challengeRequest[1:])
+ n := copy(challengeResponse[1+ed25519.PublicKeySize:], signedNonce)
+ assert.Assert(n == ed25519.SignatureSize, "copy should've copied the entire signature exactly")
+
+ _, err = conn.Write(challengeResponse)
+ if err != nil {
+ return err
+ }
+
+ 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)
+ }
+ break
+ }
+
+ if ctx.Err() != nil {
+ log.Println("context error:", ctx.Err())
+ break
+ }
+
+ err = framer.Push(ctx, buffer[:n])
+ if ctx.Err() != nil {
+ log.Println("context error:", ctx.Err())
+ break
+ }
+ assert.NoError(err, "packets from server should always be correct")
+ }
+}
+
+func handlePacketStream(framer packet.PacketFramer, program *tea.Program) {
+ for {
+ pkt, ok := <-framer.Out
+ if !ok {
+ break
+ }
+
+ payload, err := pkt.DecodedPayload()
+ assert.NoError(err, "server should always provide a decodeable packet")
+
+ if pkt.Type().IsPush() {
+ log.Println("received streamed packet:", payload)
+ 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
+ }()
+ }
+}
+
+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()
+ 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")
+ _, err := pkt.Into(conn)
+ errDeadline = conn.SetWriteDeadline(time.Time{})
+ assert.NoError(errDeadline, "setting a write deadline should not error")
+ connMu.Unlock()
+ if err != nil {
+ log.Println("request send error:", err)
+ close(responseChan)
+ return
+ }
+
+ responsesMu.Lock()
+ asyncResponses = append(asyncResponses, responseChan)
+ responsesMu.Unlock()
+
+ 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
+}
diff --git a/internal/client/server.crt b/internal/client/gateway/server.crt
index caf2384..caf2384 100644
--- a/internal/client/server.crt
+++ b/internal/client/gateway/server.crt
diff --git a/internal/client/ui.go b/internal/client/ui.go
deleted file mode 100644
index 9de322c..0000000
--- a/internal/client/ui.go
+++ /dev/null
@@ -1,125 +0,0 @@
-package client
-
-import (
- "fmt"
- "log"
- "strings"
-
- "github.com/charmbracelet/bubbles/cursor"
- "github.com/charmbracelet/bubbles/textarea"
- "github.com/charmbracelet/bubbles/viewport"
- tea "github.com/charmbracelet/bubbletea"
- "github.com/charmbracelet/lipgloss"
-
- "github.com/kyren223/eko/pkg/assert"
-)
-
-func startUI() {
- p := tea.NewProgram(initialModel(), tea.WithAltScreen())
- if _, err := p.Run(); err != nil {
- log.Println("charm ui error:", err)
- }
- p.
-}
-
-type model struct {
- viewport viewport.Model
- messages []string
- textarea textarea.Model
- senderStyle lipgloss.Style
- err error
-}
-
-func initialModel() model {
- ta := textarea.New()
- ta.Placeholder = "Send a message..."
- ta.Focus()
-
- ta.Prompt = "┃ "
- ta.CharLimit = 280
-
- ta.SetWidth(30)
- ta.SetHeight(3)
-
- // Remove cursor line styling
- ta.FocusedStyle.CursorLine = lipgloss.NewStyle()
-
- ta.ShowLineNumbers = false
-
- vp := viewport.New(30, 20)
- vp.SetContent("Welcome to Eko!\n Type a message and press Enter to send.")
-
- ta.KeyMap.InsertNewline.SetEnabled(false)
-
- messages, err := getMessages()
- assert.NoError(err, "TODO HANDLE ERROR")
- vp.SetContent(strings.Join(messages, "\n"))
-
- return model{
- textarea: ta,
- messages: messages,
- viewport: vp,
- senderStyle: lipgloss.NewStyle().Foreground(lipgloss.Color("5")),
- err: nil,
- }
-}
-
-func (m model) Init() tea.Cmd {
- return textarea.Blink
-}
-
-func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
- switch msg := msg.(type) {
- case tea.WindowSizeMsg:
- m.viewport.Width = msg.Width
- m.textarea.SetWidth(msg.Width)
- return m, tea.Printf("%dx%d", msg.Width, msg.Height)
- case tea.KeyMsg:
- switch msg.Type {
- case tea.KeyEsc, tea.KeyCtrlC:
- // Quit.
- fmt.Println(m.textarea.Value())
- return m, tea.Quit
- case tea.KeyEnter:
- value := m.textarea.Value()
-
- content := strings.TrimSpace(value)
- if content == "" {
- // Don't send empty messages.
- return m, nil
- }
- assert.NoError(sendMessage(content), "TODO HANDLE ERROR")
- messages, err := getMessages()
- assert.NoError(err, "TODO HANDLE ERROR")
-
- // m.messages = append(m.messages, m.senderStyle.Render("You: ")+content)
- m.messages = messages
- m.viewport.SetContent(strings.Join(m.messages, "\n"))
- m.textarea.Reset()
- m.viewport.GotoBottom()
- return m, nil
- default:
- // Send all other keypresses to the textarea.
- var cmd tea.Cmd
- m.textarea, cmd = m.textarea.Update(msg)
- return m, cmd
- }
-
- case cursor.BlinkMsg:
- // Textarea should also process cursor blinks.
- var cmd tea.Cmd
- m.textarea, cmd = m.textarea.Update(msg)
- return m, cmd
-
- default:
- return m, nil
- }
-}
-
-func (m model) View() string {
- return fmt.Sprintf(
- "%s\n\n%s",
- m.viewport.View(),
- m.textarea.View(),
- ) + "\n\n"
-}
diff --git a/internal/data/chat.go b/internal/data/chat.go
index 9d2443d..ae5dc67 100644
--- a/internal/data/chat.go
+++ b/internal/data/chat.go
@@ -1,6 +1,9 @@
package data
-import "github.com/kyren223/eko/pkg/snowflake"
+import (
+ "github.com/kyren223/eko/pkg/snowflake"
+ "github.com/kyren223/eko/pkg/utils"
+)
// Represents an Eko Network, equivalent to a "discord server".
type Network struct {
@@ -23,11 +26,27 @@ type Signal struct {
// Represents a message
type Message struct {
- Id snowflake.ID
- SenderId snowflake.ID
+ Id snowflake.ID
+ SenderId snowflake.ID
FrequencyId snowflake.ID
- NetworkId snowflake.ID
- Contents string
+ NetworkId snowflake.ID
+ Contents string
+}
+
+func (a Message) CmpTimestamp(b Message) int {
+ cmpTime := a.Id.Time() - b.Id.Time()
+ if cmpTime != 0 {
+ return int(utils.Clamp(cmpTime, -1, 1))
+ }
+ cmpStep := a.Id.Step() - b.Id.Step()
+ if cmpStep != 0 {
+ return int(utils.Clamp(cmpStep, -1, 1))
+ }
+ cmpNode := a.Id.Node() - b.Id.Node()
+ if cmpNode != 0 {
+ return int(utils.Clamp(cmpNode, -1, 1))
+ }
+ return 0
}
// Represents an Eko User
diff --git a/internal/packet/encoders.go b/internal/packet/encoders.go
index f5841ce..760c11c 100644
--- a/internal/packet/encoders.go
+++ b/internal/packet/encoders.go
@@ -3,11 +3,12 @@ package packet
import (
"encoding/json"
- "github.com/kyren223/eko/pkg/assert"
"github.com/vmihailenco/msgpack/v5"
+
+ "github.com/kyren223/eko/pkg/assert"
)
-type TypedMessage interface {
+type Payload interface {
Type() PacketType
}
@@ -29,24 +30,24 @@ func (e defaultPacketEncoder) Payload() []byte {
return e.data
}
-func NewJsonEncoder(message TypedMessage) PacketEncoder {
- data, err := json.Marshal(message)
+func NewJsonEncoder(payload Payload) PacketEncoder {
+ data, err := json.Marshal(payload)
assert.NoError(err, "encoding a message with JSON should never fail")
return defaultPacketEncoder{
data: data,
encoding: EncodingJson,
- packetType: message.Type(),
+ packetType: payload.Type(),
}
}
-func NewMsgPackEncoder(message TypedMessage) PacketEncoder {
- data, err := msgpack.Marshal(message)
+func NewMsgPackEncoder(payload Payload) PacketEncoder {
+ data, err := msgpack.Marshal(payload)
assert.NoError(err, "encoding a message with msg pack should never fail")
return defaultPacketEncoder{
data: data,
encoding: EncodingMsgPack,
- packetType: message.Type(),
+ packetType: payload.Type(),
}
}
diff --git a/internal/packet/packet.go b/internal/packet/packet.go
index e9d5d94..f38ad14 100644
--- a/internal/packet/packet.go
+++ b/internal/packet/packet.go
@@ -77,6 +77,15 @@ func (e PacketType) IsSupported() bool {
}
}
+func (e PacketType) IsPush() bool {
+ switch e {
+ case PacketError, PacketSendMessage:
+ return false
+ default:
+ return true
+ }
+}
+
const (
VERSION = byte(1)
PACKET_MAX_SIZE = math.MaxUint16
@@ -156,7 +165,7 @@ func (p Packet) Into(writer io.Writer) (int, error) {
return writer.Write(p.data)
}
-func (p Packet) DecodePayload(v TypedMessage) error {
+func (p Packet) DecodePayloadInto(v Payload) error {
if p.Type() != v.Type() {
return fmt.Errorf("type mismatch: want %v got %v", p.Type(), v.Type())
}
@@ -175,24 +184,40 @@ func (p Packet) DecodePayload(v TypedMessage) error {
}
}
+func (p Packet) DecodedPayload() (Payload, error) {
+ var payload Payload
+ switch p.Type() {
+ case PacketError:
+ payload = &ErrorMessage{}
+ case PacketMessages:
+ payload = &Messages{}
+ case PacketSendMessage:
+ payload = &SendMessage{}
+ default:
+ assert.Never("packet type of a packet struct must always be valid")
+ }
+ err := p.DecodePayloadInto(payload)
+ return payload, err
+}
+
var (
PacketUnsupportedVersion error = errors.New("packet error: unsupported version")
PacketUnsupportedEncoding error = errors.New("packet error: unsupported encoding")
PacketUnsupportedType error = errors.New("packet error: unsupported type")
)
-type packetFramer struct {
+type PacketFramer struct {
buffer []byte
Out chan Packet
}
-func NewFramer(ctx context.Context) packetFramer {
- return packetFramer{
+func NewFramer(ctx context.Context) PacketFramer {
+ return PacketFramer{
Out: make(chan Packet, 10),
}
}
-func (f *packetFramer) Push(ctx context.Context, data []byte) error {
+func (f *PacketFramer) Push(ctx context.Context, data []byte) error {
f.buffer = append(f.buffer, data...)
for {
@@ -208,7 +233,7 @@ func (f *packetFramer) Push(ctx context.Context, data []byte) error {
}
}
-func (f *packetFramer) parse() (*Packet, error) {
+func (f *PacketFramer) parse() (*Packet, error) {
if len(f.buffer) < HEADER_SIZE {
return nil, nil
}
diff --git a/internal/server/server.go b/internal/server/server.go
index 4ac8969..ff28ad7 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -168,6 +168,9 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
if !ok {
break
}
+ if packet.Type().IsPush() {
+ log.Println(addr, "streaming packet:", packet)
+ }
if _, err := packet.Into(conn); err != nil {
log.Println(addr, err)
break
@@ -192,7 +195,8 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
buffer := make([]byte, 512)
for {
- conn.SetReadDeadline(time.Now().Add(time.Second))
+ err := conn.SetReadDeadline(time.Now().Add(time.Second))
+ assert.NoError(err, "setting read deadline should not error")
n, err := conn.Read(buffer)
deadlineExceeded := errors.Is(err, os.ErrDeadlineExceeded)
if err != nil && !deadlineExceeded {
@@ -220,13 +224,19 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
}
func handleAuth(conn net.Conn, nonce []byte) (ed25519.PublicKey, error) {
- conn.SetDeadline(time.Now().Add(time.Second * 5))
+ 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)
- _, err := conn.Write(challengePacket)
+ _, err = conn.Write(challengePacket)
if err != nil {
return nil, fmt.Errorf("error writing challenge: %w", err)
}
@@ -252,7 +262,6 @@ func handleAuth(conn net.Conn, nonce []byte) (ed25519.PublicKey, error) {
return nil, errors.New("signature verification failed")
}
- conn.SetDeadline(time.Time{})
return pubKey, nil
}
@@ -310,20 +319,31 @@ func processPacket(ctx context.Context, pkt packet.Packet) packet.Packet {
session, ok := FromContext(ctx)
assert.Assert(ok, "context in process packet should always have a session")
- var response packet.TypedMessage
- switch pkt.Type() {
- case packet.PacketSendMessage:
- var request packet.SendMessage
- if err := pkt.DecodePayload(&request); err != nil {
- log.Println("decode error:", err)
- response = &packet.ErrorMessage{Error: "malformed payload"}
- break
- }
+ var response packet.Payload
+
+ request, err := pkt.DecodedPayload()
+ if err != nil {
+ response = &packet.ErrorMessage{Error: "malformed payload"}
+ } else {
+ response = processRequest(ctx, request)
+ }
+
+ assert.NotNil(response, "response must always be assigned to")
+ log.Println(session.Addr, "sending ", response.Type(), "response:", response)
+ return packet.NewPacket(packet.NewMsgPackEncoder(response))
+}
+
+func processRequest(ctx context.Context, request packet.Payload) packet.Payload {
+ session, ok := FromContext(ctx)
+ assert.Assert(ok, "context in process packet should always have a session")
+
+ log.Println(session.Addr, "processing", request.Type(), "request:", request)
+ switch request := request.(type) {
+ case *packet.SendMessage:
content := strings.TrimSpace(request.Content)
if content == "" {
- response = &packet.ErrorMessage{Error: "message content must not be blank"}
- break
+ return &packet.ErrorMessage{Error: "message content must not be blank"}
}
node := session.Server.Node
@@ -334,15 +354,27 @@ func processPacket(ctx context.Context, pkt packet.Packet) packet.Packet {
NetworkId: node.Generate(), // TODO: replace with actual ID
Contents: content,
}
- messages = append(messages, message)
+ sendMessage(ctx, message)
+
+ return packet.NewOkMessage()
- response = packet.NewOkMessage()
default:
- response = &packet.ErrorMessage{Error: "use of unsupported packet type"}
+ return &packet.ErrorMessage{Error: "use of unsupported packet type"}
}
-
- assert.NotNil(response, "response must always be assigned to")
- return packet.NewPacket(packet.NewMsgPackEncoder(response))
}
var messages []data.Message
+
+func sendMessage(ctx context.Context, msg data.Message) {
+ session, ok := FromContext(ctx)
+ assert.Assert(ok, "context in process packet should always have a session")
+
+ messages = append(messages, msg)
+
+ payload := &packet.Messages{
+ Messages: messages,
+ }
+
+ pkt := packet.NewPacket(packet.NewMsgPackEncoder(payload))
+ session.WriteQueue <- pkt
+}