From 755df330a4785456387ecafea2e305c1a9f4ba0a Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Thu, 10 Oct 2024 01:14:37 +0300 Subject: refactor: move to tls tcp connection for security purposes --- internal/client/client.go | 26 +++++++++--- internal/client/server.crt | 21 ++++++++++ internal/server/handler.go | 54 ++++++++++++++++++++++--- internal/server/server.crt | 21 ++++++++++ internal/server/server.go | 85 +++++++++++++++++----------------------- internal/utils/packets/packet.go | 80 ------------------------------------- 6 files changed, 147 insertions(+), 140 deletions(-) create mode 100644 internal/client/server.crt create mode 100644 internal/server/server.crt delete mode 100644 internal/utils/packets/packet.go (limited to 'internal') diff --git a/internal/client/client.go b/internal/client/client.go index b60b320..c2ef6e1 100644 --- a/internal/client/client.go +++ b/internal/client/client.go @@ -2,8 +2,10 @@ package client import ( "bufio" + "crypto/tls" + "crypto/x509" + _ "embed" "fmt" - "net" "os" "strings" "time" @@ -11,8 +13,20 @@ import ( "github.com/kyren223/eko/internal/utils/log" ) +//go:embed server.crt +var certPEM []byte + func Run() { - log.SetLevel(log.LevelDebug) + certPool := x509.NewCertPool() + if !certPool.AppendCertsFromPEM(certPEM) { + log.Fatal("failed to append server certificate") + } + + tlsConfig := &tls.Config{ + RootCAs: certPool, + ServerName: "localhost", + } + log.Info("Client started, waiting for user input...") for { fmt.Print("> ") @@ -21,20 +35,20 @@ func Run() { if input == ":q" || input == "exit" || input == "quit" { break } - log.Debug("Input: %v", input) - err := processRequest(input) + err := processRequest(input, tlsConfig) if err != nil { log.Error("%v", err) } } } -func processRequest(request string) error { - conn, err := net.Dial("tcp", ":7223") +func processRequest(request string, tlsConfig *tls.Config) error { + conn, err := tls.Dial("tcp", ":7223", tlsConfig) if err != nil { return fmt.Errorf("Unable to establish connection with server: %v", err) } defer conn.Close() + conn.SetDeadline(time.Now().Add(time.Second)) log.Info("Established connection to server: %v", conn.RemoteAddr().String()) diff --git a/internal/client/server.crt b/internal/client/server.crt new file mode 100644 index 0000000..caf2384 --- /dev/null +++ b/internal/client/server.crt @@ -0,0 +1,21 @@ +-----BEGIN CERTIFICATE----- +MIIDbTCCAlWgAwIBAgIUZyvzq7LOxKqTttTRjChoJWOO4pYwDQYJKoZIhvcNAQEL +BQAwOjEhMB8GCSqGSIb3DQEJARYSa3lyZW4yMjNAcHJvdG9uLm1lMRUwEwYDVQQD +DAxreXJlbjIyMy5kZXYwHhcNMjQxMDA5MjE1NzQzWhcNMjUxMDA5MjE1NzQzWjA6 +MSEwHwYJKoZIhvcNAQkBFhJreXJlbjIyM0Bwcm90b24ubWUxFTATBgNVBAMMDGt5 +cmVuMjIzLmRldjCCASIwDQYJKoZIhvcNAQEBBQADggEPADCCAQoCggEBAK7hd+zT +kqrn/8EhLEO0uMKKHgfoyczYWTlA9uPFADOsjdzXRLuR/Y3rK0PBE4u55xcjYZSf +mzJmVHuv1rEFOt634YOoE2UwJd9V2M0p+cD716XIEDNPfVCUe77FoZoYaH1h8QF5 +Mrx2eDH5JZt690F05O39zYzbb7+RlChWlt1kBcmLEZ1GKJeXznbL6lLMh20deYX9 +7oemqYMqP9DFbFeHkubeZ20yQvKW9cOWae9M+IhE9dAa8fm5WdfiDoTdAHfbIawx +r1OB4YqfXlXler9wAHfHWeCS0KgZCTdghF1h6wtYlwyQZcUuv+dHN7SP7zVo8pOD +b7NUqjFAMGlNgf0CAwEAAaNrMGkwCQYDVR0TBAIwADALBgNVHQ8EBAMCBaAwMAYD +VR0RBCkwJ4IJbG9jYWxob3N0ggsqLmxvY2FsaG9zdIINZWtvLmxvY2FsaG9zdDAd +BgNVHQ4EFgQUf32SXO976zgO0K/wlgWdyT3EPzcwDQYJKoZIhvcNAQELBQADggEB +AHVMGCkaZv5eIOQwevfrsEJQo3dNG34om8wBVGS5iQyho0VJZpKZSiQ16yv4x2kc +UICfVEFcfO/7/hRlA5yLWE/wpeqCgTSgtQ74gvc8D6H26wCznSPj9MIRWxYhSmPM +YO+7UKqyvFoaKiW4OkqJvCRzrpwr/lbXcGpD47UqT5gRvjJ91ULCHIUt8qDUS6+8 +mEGJAe/xFkiJ6zT0bThlqMaCA4v5g9tHGXzooIZ+YSgTvlWhAM6mVwt34l2rDSOw +4YNGUJXKCoGpy8U0NteIOOs6HhaslJpKe1mSSxmMQcgBcaf6yBT08mYfQPSsaeOk +OSoncuVBnT64liAtShpsgTc= +-----END CERTIFICATE----- diff --git a/internal/server/handler.go b/internal/server/handler.go index 9add3a0..4f99260 100644 --- a/internal/server/handler.go +++ b/internal/server/handler.go @@ -1,32 +1,74 @@ package server import ( + "fmt" "net" "sync" "github.com/kyren223/eko/internal/utils/log" + "github.com/kyren223/eko/internal/utils/packets" ) -func handleClient(conn net.Conn, wg *sync.WaitGroup) { +func handleConnection(conn net.Conn, wg *sync.WaitGroup) { log.Info("Accepted client: %v", conn.RemoteAddr().String()) defer log.Info("Disconnecting client: %v", conn.RemoteAddr().String()) defer conn.Close() defer wg.Done() + // packet, err := packets.ReadPacket(conn) + // if err != nil { + // log.Info("Received bad packet: %v", err) + // err = respondWithError(conn, err, 0) + // if err != nil { + // log.Error("Failed to respond to bad packet: %v", err) + // } + // return + // } + // + // if ok, reason := isPacketOk(packet); !ok { + // err = respondWithError(conn, fmt.Errorf("unsupported packet: %v", reason), 0) + // if err != nil { + // log.Error("Failed to respond to unsupported packet: %v", err) + // } + // return + // } + buffer := make([]byte, 1024) - n, err := conn.Read(buffer) + bytesRead, err := conn.Read(buffer) if err != nil { log.Error("Failed reading: %v", err) return } - log.Info("Read %v bytes: %v", n, string(buffer[:n])) + request := string(buffer[:bytesRead]) + log.Info("Read %v bytes: %v", bytesRead, request) - response := []byte("Server response") - n, err = conn.Write(response) + response := []byte(fmt.Sprintf("Eko \"%v\"", request)) + bytesWritten, err := conn.Write(response) if err != nil { log.Error("Failed writing response: %v", err) return } - log.Info("Written %v bytes", n) + log.Info("Written %v bytes", bytesWritten) +} + +func isPacketOk(packet packets.Packet) (bool, string) { + if packet.Version() != packets.V1 { + return false, "versions other than 1 are not supported" + } + if packet.HasFlag(packets.FlagError) { + return false, "error flag should only be used by the server" + } + return true, "" } +func respondWithError(conn net.Conn, err error, extraFlag byte) error { + flag := packets.FlagError | extraFlag + packet := packets.NewPacket(packets.V1, flag, []byte(err.Error())) + return packet.Write(conn) +} + +func processPacket(packet packets.Packet) (sessId uint16, data []byte, err error) { + if packet.HasFlag(packets.FlagHandshake) { + } + return +} diff --git a/internal/server/server.crt b/internal/server/server.crt new file mode 100644 index 0000000..caf2384 --- /dev/null +++ b/internal/server/server.crt @@ -0,0 +1,21 @@ +-----BEGIN CERTIFICATE----- +MIIDbTCCAlWgAwIBAgIUZyvzq7LOxKqTttTRjChoJWOO4pYwDQYJKoZIhvcNAQEL +BQAwOjEhMB8GCSqGSIb3DQEJARYSa3lyZW4yMjNAcHJvdG9uLm1lMRUwEwYDVQQD +DAxreXJlbjIyMy5kZXYwHhcNMjQxMDA5MjE1NzQzWhcNMjUxMDA5MjE1NzQzWjA6 +MSEwHwYJKoZIhvcNAQkBFhJreXJlbjIyM0Bwcm90b24ubWUxFTATBgNVBAMMDGt5 +cmVuMjIzLmRldjCCASIwDQYJKoZIhvcNAQEBBQADggEPADCCAQoCggEBAK7hd+zT +kqrn/8EhLEO0uMKKHgfoyczYWTlA9uPFADOsjdzXRLuR/Y3rK0PBE4u55xcjYZSf +mzJmVHuv1rEFOt634YOoE2UwJd9V2M0p+cD716XIEDNPfVCUe77FoZoYaH1h8QF5 +Mrx2eDH5JZt690F05O39zYzbb7+RlChWlt1kBcmLEZ1GKJeXznbL6lLMh20deYX9 +7oemqYMqP9DFbFeHkubeZ20yQvKW9cOWae9M+IhE9dAa8fm5WdfiDoTdAHfbIawx +r1OB4YqfXlXler9wAHfHWeCS0KgZCTdghF1h6wtYlwyQZcUuv+dHN7SP7zVo8pOD +b7NUqjFAMGlNgf0CAwEAAaNrMGkwCQYDVR0TBAIwADALBgNVHQ8EBAMCBaAwMAYD +VR0RBCkwJ4IJbG9jYWxob3N0ggsqLmxvY2FsaG9zdIINZWtvLmxvY2FsaG9zdDAd +BgNVHQ4EFgQUf32SXO976zgO0K/wlgWdyT3EPzcwDQYJKoZIhvcNAQELBQADggEB +AHVMGCkaZv5eIOQwevfrsEJQo3dNG34om8wBVGS5iQyho0VJZpKZSiQ16yv4x2kc +UICfVEFcfO/7/hRlA5yLWE/wpeqCgTSgtQ74gvc8D6H26wCznSPj9MIRWxYhSmPM +YO+7UKqyvFoaKiW4OkqJvCRzrpwr/lbXcGpD47UqT5gRvjJ91ULCHIUt8qDUS6+8 +mEGJAe/xFkiJ6zT0bThlqMaCA4v5g9tHGXzooIZ+YSgTvlWhAM6mVwt34l2rDSOw +4YNGUJXKCoGpy8U0NteIOOs6HhaslJpKe1mSSxmMQcgBcaf6yBT08mYfQPSsaeOk +OSoncuVBnT64liAtShpsgTc= +-----END CERTIFICATE----- diff --git a/internal/server/server.go b/internal/server/server.go index 8d3937d..aeb97f8 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -1,6 +1,8 @@ package server import ( + "crypto/tls" + _ "embed" "net" "os" "os/signal" @@ -11,69 +13,56 @@ import ( "github.com/kyren223/eko/internal/utils/log" ) -const PORT int = 7223 +const port = 7223 + +//go:embed server.crt +var certPEM []byte + +//go:embed server.key +var keyPEM []byte func Start() { - server, err := NewServer(PORT) + cert, err := tls.X509KeyPair(certPEM, keyPEM) if err != nil { - log.Error("Unable to start server: %v", err) - return + log.Fatal("Error loading certificate: %s", err) } - var wg sync.WaitGroup - stopChan := make(chan os.Signal, 1) - signal.Notify(stopChan, syscall.SIGINT, syscall.SIGTERM) - wg.Add(1) - go handleInterrupt(server, stopChan, &wg) + tlsConfig := &tls.Config{ + Certificates: []tls.Certificate{cert}, + } - server.Listen() - wg.Wait() -} + listener, err := tls.Listen("tcp", ":"+strconv.Itoa(port), tlsConfig) + if err != nil { + log.Fatal("Error starting listener: %s", err) + } + defer listener.Close() -func handleInterrupt(server *Server, stopChan <-chan os.Signal, wg *sync.WaitGroup) { - defer wg.Done() - <-stopChan - log.Info("Interrupt Occurred") - log.Info("Shutting down server...") - server.Close() - log.Info("Waiting for all connections to close") - server.Wait() - log.Info("Server has been shutdown") -} + signalChan := make(chan os.Signal, 1) + signal.Notify(signalChan, syscall.SIGINT, syscall.SIGTERM) + go handleInterrupt(listener, signalChan) -type Server struct { - listener net.Listener - wg sync.WaitGroup + var wg sync.WaitGroup + listen(listener, &wg) + wg.Wait() } -func NewServer(port int) (*Server, error) { - listener, err := net.Listen("tcp", ":"+strconv.Itoa(port)) - if err != nil { - return nil, err - } - - log.Info("Created server on port %v", port) - return &Server{listener, sync.WaitGroup{}}, nil +func handleInterrupt(listener net.Listener, stopChan <-chan os.Signal) { + <-stopChan + log.Info("Interrupt Signal") + log.Info("Closing listener from receiving new connections") + listener.Close() } -func (s *Server) Listen() { - log.Info("Server started listening... %v", s.listener.Addr().String()) +func listen(listener net.Listener, wg *sync.WaitGroup) { + log.Info("Started listening on port %v...", port) for { - conn, err := s.listener.Accept() + conn, err := listener.Accept() if err != nil { + log.Warn("Failed to accept connection: %v", err) break } - s.wg.Add(1) - go handleClient(conn, &s.wg) + wg.Add(1) + go handleConnection(conn, wg) } -} - -// Stop stops the server. The blocked Listen call will be unlocked -func (s *Server) Close() { - s.listener.Close() -} - -// Wait blocks until all active connections to the server are done -func (s *Server) Wait() { - s.wg.Wait() + log.Info("Stopped listening on port %v...", port) } diff --git a/internal/utils/packets/packet.go b/internal/utils/packets/packet.go deleted file mode 100644 index a83d731..0000000 --- a/internal/utils/packets/packet.go +++ /dev/null @@ -1,80 +0,0 @@ -package packets - -import ( - "bytes" - "encoding/binary" - "fmt" - "io" -) - -const magicBytes = "EKO" - -const ( - FlagHandshake byte = 0b1000_0000 - FlagV1 byte = 0b0000_0000 -) - -type Packet struct { - flag byte - data []byte -} - -func NewPacket(flag byte, data []byte) Packet { - return Packet{flag, data} -} - -func (p Packet) Write(w io.Writer) error { - if err := binary.Write(w, binary.LittleEndian, magicBytes); err != nil { - return err - } - if err := binary.Write(w, binary.LittleEndian, p.flag); err != nil { - return err - } - if err := binary.Write(w, binary.LittleEndian, uint32(len(p.data))); err != nil { - return err - } - if err := binary.Write(w, binary.LittleEndian, p.data); err != nil { - return err - } - return nil -} - -func ReadPacket(r io.Reader) (Packet, error) { - buffer := [len(magicBytes) + 1 + 4]byte{} - bytesRead, err := r.Read(buffer[:]) - if bytesRead != len(buffer) { - if err != nil && err != io.EOF { - return Packet{}, err - } - return Packet{}, fmt.Errorf("invalid packet size: got %v, want %v", bytesRead, len(buffer)) - } - - magic := buffer[:len(magicBytes)] - if !bytes.Equal(magic, []byte(magicBytes)) { - return Packet{}, fmt.Errorf("invalid magic number: got %v, want %v", magic, magicBytes) - } - - flag := buffer[len(magicBytes)] - - // TODO: Consider adding some data limit (maybe 65k?) - var length uint32 - err = binary.Read( - bytes.NewReader(buffer[len(magicBytes):len(magicBytes)+4]), - binary.LittleEndian, - &length, - ) - if err != nil { - panic(fmt.Errorf("Assertion Failed in packet.ReadPacket(io.Reader): %v", err)) - } - - data := make([]byte, length) - bytesRead, err = r.Read(data) - if uint32(bytesRead) != length { - if err != nil { - return Packet{}, err - } - return Packet{}, fmt.Errorf("invalid data size: got %v, want %v", bytesRead, length) - } - - return Packet{flag, data}, nil -} -- cgit v1.3.1