diff options
| author | Kyren223 <ulmliad223@gmail.com> | 2024-10-20 11:36:43 +0300 |
|---|---|---|
| committer | Kyren223 <ulmliad223@gmail.com> | 2024-10-20 11:36:43 +0300 |
| commit | 1997bc8a150b92783060cc7129e1ee72a761182b (patch) | |
| tree | 35d4df750feb2f1f36883ce8d7f2da9a72c59a25 /internal/server | |
| parent | 6e6ca1d7a10a3a4a1decbdb78e2ea11ce24974e3 (diff) | |
refactor: mid refactor
Diffstat (limited to 'internal/server')
| -rw-r--r-- | internal/server/handler.go | 122 | ||||
| -rw-r--r-- | internal/server/protocol.txt | 59 | ||||
| -rw-r--r-- | internal/server/server.go | 307 |
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 +) |
