summaryrefslogtreecommitdiff
path: root/internal/server
diff options
context:
space:
mode:
authorKyren223 <ulmliad223@gmail.com>2024-10-20 11:36:43 +0300
committerKyren223 <ulmliad223@gmail.com>2024-10-20 11:36:43 +0300
commit1997bc8a150b92783060cc7129e1ee72a761182b (patch)
tree35d4df750feb2f1f36883ce8d7f2da9a72c59a25 /internal/server
parent6e6ca1d7a10a3a4a1decbdb78e2ea11ce24974e3 (diff)
refactor: mid refactor
Diffstat (limited to 'internal/server')
-rw-r--r--internal/server/handler.go122
-rw-r--r--internal/server/protocol.txt59
-rw-r--r--internal/server/server.go307
3 files changed, 328 insertions, 160 deletions
diff --git a/internal/server/handler.go b/internal/server/handler.go
deleted file mode 100644
index 0ab5c33..0000000
--- a/internal/server/handler.go
+++ /dev/null
@@ -1,122 +0,0 @@
-package server
-
-import (
- "context"
- "errors"
- "fmt"
- "log"
- "net"
- "strings"
- "sync"
- "time"
-
- "github.com/kyren223/eko/internal/data"
- "github.com/kyren223/eko/internal/packet"
- "github.com/kyren223/eko/pkg/assert"
- "github.com/kyren223/eko/pkg/snowflake"
-)
-
-func handleConnection(conn net.Conn, wg *sync.WaitGroup) {
- log.Println("accepted client:", conn.RemoteAddr().String())
- defer log.Println("disconnected client:", conn.RemoteAddr().String())
- defer conn.Close()
- defer wg.Done()
-
- ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
- defer cancel()
- out, outErr := packet.RunFramer(ctx, conn)
-
-outer:
- for {
- select {
- case packet, ok := <-out:
- if !ok {
- break outer
- }
- log.Printf("client %v: request packet: %v\n", conn.RemoteAddr().String(), packet)
- responsePacket, err := handlePacket(packet)
- log.Printf("client %v: response packet: %v\n", conn.RemoteAddr().String(), responsePacket)
- if err != nil {
- log.Printf("client %v: error processing request: %v\n", conn.RemoteAddr().String(), err)
- break outer
- }
- err = responsePacket.Into(conn)
- if err != nil {
- log.Printf("client %v: error writing packet: %v\n", conn.RemoteAddr().String(), err)
- break outer
- }
-
- case err := <-outErr:
- if err == packet.PacketUnsupportedEncoding {
- err := unsupportedEncodingErrorPacket.Into(conn)
- log.Printf("client %v: error writing unsupported encoding packet: %v\n", conn.RemoteAddr().String(), err)
- } else if err == packet.PacketUnsupportedType {
- err := unsupportedTypeErrorPacket.Into(conn)
- log.Printf("client %v: error writing unsupported type packet: %v\n", conn.RemoteAddr().String(), err)
- } else if err != nil {
- log.Printf("client %v: internal error: %v\n", conn.RemoteAddr().String(), err)
- }
- break outer
-
- case <-ctx.Done():
- log.Printf("client %v: %v\n", conn.RemoteAddr().String(), ctx.Err())
- break outer
- }
- }
-}
-
-func handlePacket(pkt packet.Packet) (packet.Packet, error) {
- var response packet.TypedMessage
- switch pkt.Type() {
- case packet.TypeEko:
- var request packet.EkoMessage
- if err := pkt.DecodePayload(&request); err != nil {
- return packet.Packet{}, fmt.Errorf("decode error: %v", err)
- }
-
- response = &packet.EkoMessage{Message: "Eko \"" + request.Message + "\""}
- case packet.TypeSendMessage:
- var request packet.SendMessageMessage
- if err := pkt.DecodePayload(&request); err != nil {
- return packet.Packet{}, fmt.Errorf("decode error: %v", err)
- }
-
- content := strings.TrimSpace(request.Content)
- if content == "" {
- response = &packet.ErrorMessage{Error: "content must not be blank"}
- break
- }
-
- message := data.Message{
- Id: node.Generate(),
- SenderId: node.Generate(),
- FrequencyId: node.Generate(),
- NetworkId: node.Generate(),
- Contents: content,
- }
- messages = append(messages, message)
-
- response = &packet.EkoMessage{Message: "Eko OK"}
- case packet.TypeGetMessages:
- var request packet.GetMessagesMessage
- if err := pkt.DecodePayload(&request); err != nil {
- return packet.Packet{}, fmt.Errorf("decode error: %v", err)
- }
-
- response = &packet.MessagesMessage{Messages: messages}
- default:
- return packet.Packet{}, errors.New("TODO: not implemented yet")
- }
-
- assert.NotNil(response, "response must always be set")
- encoder, err := packet.NewMsgPackEncoder(response)
- if err != nil {
- return packet.Packet{}, fmt.Errorf("encode error: %v", err)
- }
- return packet.NewPacket(encoder), nil
-}
-
-var (
- node = snowflake.NewNode(1)
- messages []data.Message
-)
diff --git a/internal/server/protocol.txt b/internal/server/protocol.txt
new file mode 100644
index 0000000..9a0070d
--- /dev/null
+++ b/internal/server/protocol.txt
@@ -0,0 +1,59 @@
+# Eko Protocol V1
+
+## Packet Structure
+
+ 0 1 2 3
+ 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1
++-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
+| Version |En.| Type | Payload Length |
++-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
+| Payload... Payload Length bytes ... |
++-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
+
+Order of bytes is from left to right, top to bottom.
+The first byte is always the version, any bytes after it
+depend on the specific value of the first byte.
+
+- Encoding: 0-3, determines the way the payload was encoded
+ - 0: JSON
+ - 1: MsgPack
+ - 2: Reserved for future use
+ - 3: Reserved for future use
+- Type: 0-63, determines the type ("schema"), of the payload
+- Payload Length: 0-65531, determines how long the payload is in bytes
+- Payload: 0 to 65531 bytes long, depending on the payload size (~64kb)
+
+## Handshake
+
+The first time a connection is established, the following packets are exchanged.
+
+- Server sends a special 1-byte for version then 32-byte the challenge nonce packet
+- Client sends a Challenge Response packet with:
+ * version (1-byte long)
+ - the client's ed25519 public key (32-bytes long)
+ - the client's ed25519 signature for the server-given nonce (64-bytes long)
+
+After the handshake the server may close the connection,
+for example due to an invalid signature.
+
+## Error handling
+
+The server may abruptly close a connection in these cases:
+
+- After the initial handshake
+- After any response
+ A server may not close the connection if it received a request, it must first response then close.
+
+The client may abruptly close a connection at any time
+
+### Malformed Packets
+
+- unsupported/invalid version: connection can be closed immediately
+- unsupported encoding: server must respond with an error type, may use any encoding, client may close the connection
+- unknown type: server must respond with an error, client may close the connection
+- malformed paylod: server must respond with an error, client may close the connection
+
+For application errors such as a client asking to send a message in a non-existent Frequency,
+the server must respond with an error packet.
+For internal errors such as database failure, the server must respond, it may choose to
+disclose as much information as it wants, or just say "internal server error".
diff --git a/internal/server/server.go b/internal/server/server.go
index 89434d3..933ad29 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -1,64 +1,92 @@
package server
import (
+ "context"
+ "crypto/rand"
"crypto/tls"
_ "embed"
"errors"
+ "fmt"
+ "io"
"log"
"net"
- "os"
- "os/signal"
"strconv"
+ "strings"
"sync"
- "syscall"
+ "time"
+ "github.com/kyren223/eko/internal/data"
"github.com/kyren223/eko/internal/packet"
"github.com/kyren223/eko/pkg/assert"
+ "github.com/kyren223/eko/pkg/snowflake"
)
-const port = 7223
-
//go:embed server.crt
var certPEM []byte
//go:embed server.key
var keyPEM []byte
-func Start() {
+var (
+ nodeId int64 = 0
+ tlsConfig *tls.Config
+
+ unsupportedEncodingErrorPacket packet.Packet
+ unsupportedTypeErrorPacket packet.Packet
+)
+
+var ErrClosedNilListener error = errors.New("server: close on nil listener")
+
+func init() {
cert, err := tls.X509KeyPair(certPEM, keyPEM)
if err != nil {
log.Fatalln("error loading certificate:", err)
}
- tlsConfig := &tls.Config{
+ tlsConfig = &tls.Config{
Certificates: []tls.Certificate{cert},
}
- prepareConstPackets()
+ message := packet.ErrorMessage{Error: packet.PacketUnsupportedEncoding.Error()}
+ encoder, err := packet.NewMsgPackEncoder(&message)
+ assert.NoError(err, "constant packets should not error")
+ unsupportedEncodingErrorPacket = packet.NewPacket(encoder)
+
+ message = packet.ErrorMessage{Error: packet.PacketUnsupportedType.Error()}
+ encoder, err = packet.NewMsgPackEncoder(&message)
+ assert.NoError(err, "constant packets should not error")
+ unsupportedTypeErrorPacket = packet.NewPacket(encoder)
+}
+
+type server struct {
+ node *snowflake.Node
+ port uint16
+}
- listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(port), tlsConfig)
+func NewServer(port uint16) server {
+ assert.Assert(nodeId <= snowflake.NodeMax, "maximum amount of servers reached %v", snowflake.NodeMax)
+ node := snowflake.NewNode(nodeId)
+ nodeId++
+
+ return server{
+ node: node,
+ port: port,
+ }
+}
+
+func (s *server) ListenAndServe(ctx context.Context) {
+ listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(int(s.port)), tlsConfig)
if err != nil {
log.Fatalf("error starting server: %s", err)
}
defer listener.Close()
+ go func() {
+ <-ctx.Done()
+ listener.Close()
+ }()
- signalChan := make(chan os.Signal, 1)
- signal.Notify(signalChan, syscall.SIGINT, syscall.SIGTERM)
- go handleInterrupt(listener, signalChan)
-
+ log.Printf("started listening on port %v...\n", s.port)
var wg sync.WaitGroup
- listen(listener, &wg)
- wg.Wait()
-}
-
-func handleInterrupt(listener net.Listener, stopChan <-chan os.Signal) {
- signal := <-stopChan
- log.Println("signal:", signal.String())
- listener.Close()
-}
-
-func listen(listener net.Listener, wg *sync.WaitGroup) {
- log.Printf("started listening on port %v...\n", port)
for {
conn, err := listener.Accept()
if err != nil {
@@ -68,22 +96,225 @@ func listen(listener net.Listener, wg *sync.WaitGroup) {
break
}
wg.Add(1)
- go handleConnection(conn, wg)
+ go func() {
+ handleConnection(ctx, conn)
+ wg.Done()
+ }()
}
- log.Printf("stopped listening on port %v...\n", port)
+ log.Printf("stopped listening on port %v\n", s.port)
+
+ log.Println("waiting for all active connections to close...")
+ wg.Wait()
+ log.Println("server shutdown complete")
}
-var unsupportedEncodingErrorPacket packet.Packet
-var unsupportedTypeErrorPacket packet.Packet
+func handleConnection(ctx context.Context, conn net.Conn) {
+ addr, ok := conn.RemoteAddr().(*net.TCPAddr)
+ assert.Assert(ok, "getting tcp address should never fail as we are using tcp connections")
-func prepareConstPackets() {
- message := packet.ErrorMessage{Error: packet.PacketUnsupportedEncoding.Error()}
- encoder, err := packet.NewMsgPackEncoder(&message)
- assert.NoError(err, "constant packets should not error")
- unsupportedEncodingErrorPacket = packet.NewPacket(encoder)
+ writeQueue := make(chan packet.Packet, 10)
+ session := newSession(addr, writeQueue)
+ nonce := session.Challenge()
- message = packet.ErrorMessage{Error: packet.PacketUnsupportedType.Error()}
- encoder, err = packet.NewMsgPackEncoder(&message)
- assert.NoError(err, "constant packets should not error")
- unsupportedTypeErrorPacket = packet.NewPacket(encoder)
+ ctx = newContext(ctx, session)
+ framer := packet.NewFramer(ctx)
+
+ log.Println(addr, "accepted")
+
+ defer func() {
+ conn.Close()
+ log.Println(addr, "disconnected")
+ }()
+
+ go func() {
+ var mu sync.Mutex
+ for {
+ packet, ok := <-writeQueue
+ if !ok {
+ break
+ }
+ mu.Lock()
+ packet.Into(conn)
+ mu.Unlock()
+ }
+ }()
+
+ go func() {
+ for {
+ request, ok := <-framer.Out
+ if !ok {
+ break
+ }
+ response := processPacket(ctx, request)
+ writeQueue <- response
+ }
+ }()
+
+ buffer := make([]byte, 512)
+ for {
+ n, err := conn.Read(buffer)
+ if err != nil {
+ if !errors.Is(err, io.EOF) {
+ log.Println(addr, "read error:", err)
+ }
+ break
+ }
+
+ err = framer.Push(ctx, buffer[:n])
+ if err != nil {
+ // Wrap err and send to client then break
+ }
+ }
+}
+
+func _handleConnection(ctx context.Context, conn net.Conn) {
+ addr, ok := conn.RemoteAddr().(*net.TCPAddr)
+ assert.Assert(ok, "getting tcp address should be valid")
+
+ log.Println(addr, "accepted")
+ defer log.Println(addr, "disconnected")
+ defer conn.Close()
+
+ // ctx = newContext(ctx, addr)
+ // // TODO: consider adding timeout/deadline to ctx?
+
+ out, outErr := packet.RunFramer(ctx, conn)
+outer:
+ for {
+ select {
+ case packet, ok := <-out:
+ if !ok {
+ break outer
+ }
+ log.Printf("client %v: request packet: %v\n", conn.RemoteAddr().String(), packet)
+ responsePacket, err := handlePacket(packet)
+ log.Printf("client %v: response packet: %v\n", conn.RemoteAddr().String(), responsePacket)
+ if err != nil {
+ log.Printf("client %v: error processing request: %v\n", conn.RemoteAddr().String(), err)
+ break outer
+ }
+ _, err = responsePacket.Into(conn)
+ if err != nil {
+ log.Printf("client %v: error writing packet: %v\n", conn.RemoteAddr().String(), err)
+ break outer
+ }
+
+ case err := <-outErr:
+ if err == packet.PacketUnsupportedEncoding {
+ _, err := unsupportedEncodingErrorPacket.Into(conn)
+ log.Printf("client %v: error writing unsupported encoding packet: %v\n", conn.RemoteAddr().String(), err)
+ } else if err == packet.PacketUnsupportedType {
+ _, err := unsupportedTypeErrorPacket.Into(conn)
+ log.Printf("client %v: error writing unsupported type packet: %v\n", conn.RemoteAddr().String(), err)
+ } else if err != nil {
+ log.Printf("client %v: internal error: %v\n", conn.RemoteAddr().String(), err)
+ }
+ break outer
+
+ case <-ctx.Done():
+ log.Printf("client %v: %v\n", conn.RemoteAddr().String(), ctx.Err())
+ break outer
+ }
+ }
+}
+
+type Session struct {
+ Addr *net.TCPAddr
+ WriteQueue <-chan packet.Packet
+
+ mu sync.Mutex
+ challenge []byte
+ issuedTime time.Time
+}
+
+func newSession(addr *net.TCPAddr, writeQueue <-chan packet.Packet) *Session {
+ session := &Session{
+ Addr: addr,
+ WriteQueue: writeQueue,
+ challenge: make([]byte, 32), // Recommended nonce size
+ }
+ session.Challenge() // Make sure an initial nonce is generated
+ return session
+}
+
+func (s *Session) Challenge() []byte {
+ s.mu.Lock()
+ defer s.mu.Unlock()
+ if time.Since(s.issuedTime) > time.Minute {
+ s.issuedTime = time.Now()
+ _, err := rand.Read(s.challenge)
+ assert.NoError(err, "random should always produce a value")
+ }
+ return s.challenge
}
+
+type key struct{}
+
+var sessKey key
+
+func newContext(ctx context.Context, sess *Session) context.Context {
+ return context.WithValue(ctx, sessKey, sess)
+}
+
+func FromContext(ctx context.Context) (*Session, bool) {
+ sess, ok := ctx.Value(sessKey).(*Session)
+ return sess, ok
+}
+
+// TODO: Move everything below this to somewhere else
+
+func handlePacket(pkt packet.Packet) (packet.Packet, error) {
+ var response packet.TypedMessage
+ switch pkt.Type() {
+ case packet.PacketEko:
+ var request packet.EkoMessage
+ if err := pkt.DecodePayload(&request); err != nil {
+ return packet.Packet{}, fmt.Errorf("decode error: %v", err)
+ }
+
+ response = &packet.EkoMessage{Message: "Eko \"" + request.Message + "\""}
+ case packet.PacketSendMessage:
+ var request packet.SendMessageMessage
+ if err := pkt.DecodePayload(&request); err != nil {
+ return packet.Packet{}, fmt.Errorf("decode error: %v", err)
+ }
+
+ content := strings.TrimSpace(request.Content)
+ if content == "" {
+ response = &packet.ErrorMessage{Error: "content must not be blank"}
+ break
+ }
+
+ message := data.Message{
+ Id: node.Generate(),
+ SenderId: node.Generate(),
+ FrequencyId: node.Generate(),
+ NetworkId: node.Generate(),
+ Contents: content,
+ }
+ messages = append(messages, message)
+
+ response = &packet.EkoMessage{Message: "Eko OK"}
+ case packet.PacketGetMessages:
+ var request packet.GetMessagesMessage
+ if err := pkt.DecodePayload(&request); err != nil {
+ return packet.Packet{}, fmt.Errorf("decode error: %v", err)
+ }
+
+ response = &packet.MessagesMessage{Messages: messages}
+ default:
+ return packet.Packet{}, errors.New("TODO: not implemented yet")
+ }
+
+ assert.NotNil(response, "response must always be set")
+ encoder, err := packet.NewMsgPackEncoder(response)
+ if err != nil {
+ return packet.Packet{}, fmt.Errorf("encode error: %v", err)
+ }
+ return packet.NewPacket(encoder), nil
+}
+
+var (
+ node = snowflake.NewNode(1)
+ messages []data.Message
+)