summaryrefslogtreecommitdiff
path: root/internal/server
diff options
context:
space:
mode:
Diffstat (limited to 'internal/server')
-rw-r--r--internal/server/api/api.go34
-rw-r--r--internal/server/server.go19
-rw-r--r--internal/server/session/session.go7
3 files changed, 52 insertions, 8 deletions
diff --git a/internal/server/api/api.go b/internal/server/api/api.go
index 3d4d274..f932fbb 100644
--- a/internal/server/api/api.go
+++ b/internal/server/api/api.go
@@ -2,13 +2,17 @@ package api
import (
"context"
+ "crypto/ed25519"
+ "database/sql"
"log"
+ "strconv"
"strings"
"github.com/kyren223/eko/internal/data"
"github.com/kyren223/eko/internal/packet"
"github.com/kyren223/eko/internal/server/session"
"github.com/kyren223/eko/pkg/assert"
+ "github.com/kyren223/eko/pkg/snowflake"
)
func SendMessage(ctx context.Context, request *packet.SendMessage) packet.Payload {
@@ -51,3 +55,33 @@ func GetMessages(ctx context.Context, request *packet.GetMessagesRange) packet.P
}
return &packet.Messages{Messages: messages}
}
+
+func GetUserById(ctx context.Context, request *packet.GetUserByID) packet.Payload {
+ queries := data.New(db)
+ user, err := queries.GetUserById(ctx, request.UserID)
+ if err == sql.ErrNoRows {
+ return &packet.Users{Users: []data.User{}}
+ }
+ if err != nil {
+ log.Println("database error when retrieving user by id:", err)
+ return &packet.ErrorMessage{Error: "internal server error"}
+ }
+ return &packet.Users{Users: []data.User{user}}
+}
+
+func CreateOrGetUser(ctx context.Context, node *snowflake.Node, pubKey ed25519.PublicKey) (data.User, error) {
+ queries := data.New(db)
+ user, err := queries.GetUserByPublicKey(ctx, pubKey)
+ if err == sql.ErrNoRows {
+ id := node.Generate()
+ user, err = queries.CreateUser(ctx, data.CreateUserParams{
+ ID: id,
+ Name: "User" + strconv.FormatInt(id.Time() % 1000, 10),
+ PublicKey: pubKey,
+ })
+ }
+ if err != nil {
+ return data.User{}, err
+ }
+ return user, nil
+}
diff --git a/internal/server/server.go b/internal/server/server.go
index 09f46e4..07cd537 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -147,9 +147,14 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
return
}
- // TODO: replace this with DB query for id
- id := server.Node().Generate()
- sess := session.NewSession(server, addr, id, pubKey)
+ user, err := api.CreateOrGetUser(ctx, server.Node(), pubKey)
+ if err != nil {
+ log.Println(addr, "user creation/fetching error:", err)
+ conn.Close()
+ log.Println(addr, "disconnected")
+ return
+ }
+ sess := session.NewSession(server, addr, user.ID, pubKey)
server.AddSession(sess)
ctx = session.NewContext(ctx, sess)
framer := packet.NewFramer()
@@ -290,11 +295,15 @@ func processRequest(ctx context.Context, request packet.Payload) packet.Payload
assert.Assert(ok, "context in process packet should always have a session")
log.Println(session.Addr(), "processing", request.Type(), "request:", request)
+ // TODO: add a way to measure the time each request/response took and log it
+ // Potentially even separate time for code vs DB operations
switch request := request.(type) {
case *packet.SendMessage:
- return timeout(50 * time.Millisecond, api.SendMessage, ctx, request)
+ return timeout(20*time.Millisecond, api.SendMessage, ctx, request)
case *packet.GetMessagesRange:
- return timeout(100 * time.Millisecond, api.GetMessages, ctx, request)
+ return timeout(50*time.Millisecond, api.GetMessages, ctx, request)
+ case *packet.GetUserByID:
+ return timeout(50*time.Millisecond, api.GetUserById, ctx, request)
default:
return &packet.ErrorMessage{Error: "use of disallowed packet type for request"}
}
diff --git a/internal/server/session/session.go b/internal/server/session/session.go
index ff56780..863b3c3 100644
--- a/internal/server/session/session.go
+++ b/internal/server/session/session.go
@@ -38,11 +38,12 @@ type Session struct {
func NewSession(manager SessionManager, addr *net.TCPAddr, id snowflake.ID, pubKey ed25519.PublicKey) *Session {
session := &Session{
+ WriteQueue: make(chan packet.Packet, 10),
+ PubKey: pubKey,
manager: manager,
addr: addr,
- PubKey: pubKey,
- WriteQueue: make(chan packet.Packet, 10),
- challenge: make([]byte, 32), // Recommended nonce size
+ id: id,
+ challenge: make([]byte, 32),
}
session.Challenge() // Make sure an initial nonce is generated
return session