summaryrefslogtreecommitdiff
path: root/internal/server/server.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/server/server.go')
-rw-r--r--internal/server/server.go148
1 files changed, 31 insertions, 117 deletions
diff --git a/internal/server/server.go b/internal/server/server.go
index ff28ad7..ab178eb 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -13,28 +13,25 @@ import (
"net"
"os"
"strconv"
- "strings"
"sync"
"time"
- "github.com/kyren223/eko/internal/data"
"github.com/kyren223/eko/internal/packet"
+ "github.com/kyren223/eko/internal/server/api"
+ "github.com/kyren223/eko/internal/server/session"
"github.com/kyren223/eko/pkg/assert"
"github.com/kyren223/eko/pkg/snowflake"
)
-//go:embed server.crt
+//go:embed certs/server.crt
var certPEM []byte
-//go:embed server.key
+//go:embed certs/server.key
var keyPEM []byte
var (
nodeId int64 = 0
tlsConfig *tls.Config
-
- unsupportedEncodingErrorPacket packet.Packet
- unsupportedTypeErrorPacket packet.Packet
)
var ErrClosedNilListener error = errors.New("server: close on nil listener")
@@ -48,18 +45,12 @@ func init() {
tlsConfig = &tls.Config{
Certificates: []tls.Certificate{cert},
}
-
- message := packet.ErrorMessage{Error: packet.PacketUnsupportedEncoding.Error()}
- unsupportedEncodingErrorPacket = packet.NewPacket(packet.NewMsgPackEncoder(&message))
-
- message = packet.ErrorMessage{Error: packet.PacketUnsupportedType.Error()}
- unsupportedTypeErrorPacket = packet.NewPacket(packet.NewMsgPackEncoder(&message))
}
type server struct {
- Node *snowflake.Node
+ node *snowflake.Node
Port uint16
- sessions map[snowflake.ID]*Session
+ sessions map[snowflake.ID]*session.Session
sessMu sync.RWMutex
}
@@ -69,16 +60,16 @@ func NewServer(port uint16) server {
nodeId++
return server{
- Node: node,
+ node: node,
Port: port,
- sessions: map[snowflake.ID]*Session{},
+ sessions: map[snowflake.ID]*session.Session{},
}
}
-func (s *server) AddSession(session *Session) {
+func (s *server) AddSession(session *session.Session) {
s.sessMu.Lock()
defer s.sessMu.Unlock()
- s.sessions[session.ID] = session
+ s.sessions[session.ID()] = session
}
func (s *server) RemoveSession(id snowflake.ID) {
@@ -87,13 +78,17 @@ func (s *server) RemoveSession(id snowflake.ID) {
delete(s.sessions, id)
}
-func (s *server) Session(id snowflake.ID) (*Session, bool) {
+func (s *server) Session(id snowflake.ID) (*session.Session, bool) {
s.sessMu.RLock()
defer s.sessMu.RUnlock()
session, ok := s.sessions[id]
return session, ok
}
+func (s *server) Node() *snowflake.Node {
+ return s.node
+}
+
func (s *server) ListenAndServe(ctx context.Context) {
listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(int(s.Port)), tlsConfig)
if err != nil {
@@ -148,23 +143,24 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
}
// TODO: replace this with DB query for id
- id := server.Node.Generate()
- session := newSession(server, addr, id, pubKey)
- server.AddSession(session)
- ctx = newContext(ctx, session)
+ id := server.Node().Generate()
+ sess := session.NewSession(server, addr, id, pubKey)
+ server.AddSession(sess)
+ ctx = session.NewContext(ctx, sess)
framer := packet.NewFramer(ctx)
defer func() {
conn.Close()
close(framer.Out)
- server.RemoveSession(session.ID)
- close(session.WriteQueue)
+ server.RemoveSession(sess.ID())
+ close(sess.WriteQueue)
+ sess.WriteQueue = nil
log.Println(addr, "disconnected")
}()
go func() {
for {
- packet, ok := <-session.WriteQueue
+ packet, ok := <-sess.WriteQueue
if !ok {
break
}
@@ -176,7 +172,6 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
break
}
}
- session.WriteQueue = nil
}()
go func() {
@@ -186,10 +181,10 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
break
}
response := processPacket(ctx, request)
- if session.WriteQueue == nil {
+ if sess.WriteQueue == nil {
break
}
- session.WriteQueue <- response
+ sess.WriteQueue <- response
}
}()
@@ -216,7 +211,9 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
if ctx.Err() != nil {
log.Println(addr, ctx.Err())
} else {
- // TODO: Wrap err and send to client then break
+ payload := packet.ErrorMessage{Error: err.Error()}
+ pkt := packet.NewPacket(packet.NewMsgPackEncoder(&payload))
+ sess.WriteQueue <- pkt
}
break
}
@@ -265,58 +262,8 @@ func handleAuth(conn net.Conn, nonce []byte) (ed25519.PublicKey, error) {
return pubKey, nil
}
-type Session struct {
- Server *server
- Addr *net.TCPAddr
- WriteQueue chan packet.Packet
- ID snowflake.ID
- PubKey ed25519.PublicKey
-
- mu sync.Mutex
- challenge []byte
- issuedTime time.Time
-}
-
-func newSession(server *server, addr *net.TCPAddr, id snowflake.ID, pubKey ed25519.PublicKey) *Session {
- session := &Session{
- Server: server,
- Addr: addr,
- PubKey: pubKey,
- WriteQueue: make(chan packet.Packet, 10),
- 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 processPacket(ctx context.Context, pkt packet.Packet) packet.Packet {
- session, ok := FromContext(ctx)
+ session, ok := session.FromContext(ctx)
assert.Assert(ok, "context in process packet should always have a session")
var response packet.Payload
@@ -334,47 +281,14 @@ func processPacket(ctx context.Context, pkt packet.Packet) packet.Packet {
}
func processRequest(ctx context.Context, request packet.Payload) packet.Payload {
- session, ok := FromContext(ctx)
+ session, ok := session.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 == "" {
- return &packet.ErrorMessage{Error: "message content must not be blank"}
- }
-
- node := session.Server.Node
- message := data.Message{
- Id: node.Generate(),
- SenderId: session.ID,
- FrequencyId: node.Generate(), // TODO: replace with actual ID
- NetworkId: node.Generate(), // TODO: replace with actual ID
- Contents: content,
- }
- sendMessage(ctx, message)
-
- return packet.NewOkMessage()
-
+ return api.SendMessage(ctx, request)
default:
return &packet.ErrorMessage{Error: "use of unsupported packet type"}
}
}
-
-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
-}