diff options
Diffstat (limited to 'internal/server')
| -rw-r--r-- | internal/server/api/api.go | 34 | ||||
| -rw-r--r-- | internal/server/server.go | 19 | ||||
| -rw-r--r-- | internal/server/session/session.go | 7 |
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 |
