summaryrefslogtreecommitdiff
path: root/internal/client/client.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/client/client.go')
-rw-r--r--internal/client/client.go98
1 files changed, 60 insertions, 38 deletions
diff --git a/internal/client/client.go b/internal/client/client.go
index 83b3d90..99c53b3 100644
--- a/internal/client/client.go
+++ b/internal/client/client.go
@@ -1,79 +1,94 @@
package client
import (
- "bufio"
"context"
"crypto/tls"
"crypto/x509"
_ "embed"
"fmt"
"log"
- "net"
"os"
- "strings"
"time"
"github.com/kyren223/eko/internal/packet"
+ "github.com/kyren223/eko/pkg/assert"
)
//go:embed server.crt
var certPEM []byte
+var tlsConfig *tls.Config
+
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)
+
certPool := x509.NewCertPool()
if !certPool.AppendCertsFromPEM(certPEM) {
log.Fatalln("failed to append server certificate")
}
- tlsConfig := &tls.Config{
+ tlsConfig = &tls.Config{
RootCAs: certPool,
ServerName: "localhost",
}
log.Println("client started, waiting for user input...")
- for {
- fmt.Print("> ")
- input, _ := bufio.NewReader(os.Stdin).ReadString('\n')
- input = strings.TrimSpace(input)
- if input == ":q" || input == "exit" || input == "quit" {
- break
- }
- err := processRequest(input, tlsConfig)
- if err != nil {
- log.Println(err)
- }
+ 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)
+ // }
+ // }
+}
+
+func sendMessage(message string) error {
+ request := packet.SendMessageMessage{Content: message}
+ var response packet.EkoMessage
+ if err := SendAndReceive(&request, &response); err != nil {
+ return err
+ }
+ 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 processRequest(input string, tlsConfig *tls.Config) error {
+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())
- if input == "SHOW" {
- request := packet.GetMessagesMessage{}
- var response packet.MessagesMessage
- if err := SendAndReceive(conn, &request, &response); err != nil {
- return err
- }
- log.Println("server response:", response.Messages)
- } else {
- request := packet.SendMessageMessage{Content: input}
- var response packet.EkoMessage
- if err := SendAndReceive(conn, &request, &response); err != nil {
- return err
- }
- log.Println("server response:", response.Message)
- }
-
- return nil
-}
-
-func SendAndReceive(conn net.Conn, request packet.TypedMessage, response packet.TypedMessage) error {
encoder, err := packet.NewMsgPackEncoder(request)
if err != nil {
return fmt.Errorf("error encoding request: %v", err)
@@ -92,7 +107,14 @@ func SendAndReceive(conn net.Conn, request packet.TypedMessage, response packet.
select {
case responsePacket := <-out:
if err := responsePacket.DecodePayload(response); err != nil {
- return fmt.Errorf("error decoding response: %v", err)
+ if responsePacket.Type() != packet.TypeError {
+ return fmt.Errorf("error decoding response: %v", err)
+ }
+ var errorResponse packet.ErrorMessage
+ if err := responsePacket.DecodePayload(&errorResponse); err != nil {
+ return fmt.Errorf("error decoding error packet: %w", err)
+ }
+ return fmt.Errorf("server error: %v", errorResponse.Error)
}
case err := <-outErr: