summaryrefslogtreecommitdiff
path: root/internal/server
diff options
context:
space:
mode:
Diffstat (limited to 'internal/server')
-rw-r--r--internal/server/api/api.go93
-rw-r--r--internal/server/api/helpers.go23
-rw-r--r--internal/server/api/models.go15
-rw-r--r--internal/server/server.go3
4 files changed, 133 insertions, 1 deletions
diff --git a/internal/server/api/api.go b/internal/server/api/api.go
index 6c8d51a..f320838 100644
--- a/internal/server/api/api.go
+++ b/internal/server/api/api.go
@@ -4,6 +4,7 @@ import (
"context"
"crypto/ed25519"
"database/sql"
+ "fmt"
"log"
"strconv"
"strings"
@@ -14,6 +15,8 @@ import (
"github.com/kyren223/eko/pkg/snowflake"
)
+var internalError = packet.Error{Error: "internal server error"}
+
func SendMessage(ctx context.Context, sess *session.Session, request *packet.SendMessage) packet.Payload {
if (request.ReceiverID != nil) == (request.FrequencyID != nil) {
return &packet.Error{Error: "either receiver id or frequency id must exist"}
@@ -34,7 +37,7 @@ func SendMessage(ctx context.Context, sess *session.Session, request *packet.Sen
})
if err != nil {
log.Println(sess.Addr(), "database error:", err, "in SendMessage")
- return &packet.Error{Error: "internal server error"}
+ return &internalError
}
return &packet.MessagesInfo{Messages: []data.Message{message}}
@@ -79,3 +82,91 @@ func CreateOrGetUser(ctx context.Context, node *snowflake.Node, pubKey ed25519.P
}
return user, nil
}
+
+func CreateNetwork(ctx context.Context, sess *session.Session, request packet.CreateNetwork) packet.Payload {
+ name := strings.TrimSpace(request.Name)
+ if name == "" {
+ return &packet.Error{Error: "server name must not be blank"}
+ }
+
+ if len(request.Icon) > MaxIconSize {
+ return &packet.Error{Error: fmt.Sprintf("icon is too large, must be smaller than %v bytes", MaxIconSize)}
+ }
+
+ if ok, err := isValidHexColor(request.BgHexColor); !ok {
+ return &packet.Error{Error: err}
+ }
+ if ok, err := isValidHexColor(request.FgHexColor); !ok {
+ return &packet.Error{Error: err}
+ }
+
+ tx, err := db.BeginTx(ctx, nil)
+ if err != nil {
+ log.Println("database error:", err)
+ return &internalError
+ }
+
+ queries := data.New(db)
+ qtx := queries.WithTx(tx)
+
+ network, err := qtx.CreateNetwork(ctx, data.CreateNetworkParams{
+ ID: sess.Manager().Node().Generate(),
+ OwnerID: sess.ID(),
+ Name: name,
+ IsPublic: request.IsPublic,
+ Icon: request.Icon,
+ BgHexColor: request.BgHexColor,
+ FgHexColor: request.FgHexColor,
+ })
+ if err != nil {
+ log.Println("database error:", err)
+ return &internalError
+ }
+
+ frequency, err := qtx.CreateFrequency(ctx, data.CreateFrequencyParams{
+ ID: sess.Manager().Node().Generate(),
+ NetworkID: network.ID,
+ Name: DefaultFrequencyName,
+ HexColor: nil,
+ Perms: PermReadWrite,
+ })
+ if err != nil {
+ log.Println("database error:", err)
+ return &internalError
+ }
+
+ networkUser, err := qtx.SetNetworkUser(ctx, data.SetNetworkUserParams{
+ UserID: network.OwnerID,
+ NetworkID: network.ID,
+ IsMember: true,
+ IsAdmin: true,
+ IsMuted: false,
+ IsBanned: false,
+ BanReason: nil,
+ })
+ if err != nil {
+ log.Println("database error:", err)
+ return &internalError
+ }
+
+ user, err := qtx.GetUserById(ctx, network.OwnerID)
+ if err != nil {
+ log.Println("database error:", err)
+ return &internalError
+ }
+
+ fullNetwork := packet.FullNetwork{
+ Network: network,
+ Frequencies: []data.Frequency{frequency},
+ Members: []packet.Member{{
+ JoinedAt: networkUser.JoinedAt,
+ User: user,
+ IsAdmin: networkUser.IsAdmin,
+ IsMuted: networkUser.IsMuted,
+ }},
+ }
+ return &packet.NetworksInfo{
+ Networks: []packet.FullNetwork{fullNetwork},
+ }
+}
+
diff --git a/internal/server/api/helpers.go b/internal/server/api/helpers.go
new file mode 100644
index 0000000..967f308
--- /dev/null
+++ b/internal/server/api/helpers.go
@@ -0,0 +1,23 @@
+package api
+
+import "strings"
+
+const hex = "0123456789abcdefABCDEF"
+
+func isValidHexColor(color string) (bool, string) {
+ if len(color) != 7 {
+ return false, "color must be hex with length of 7"
+ }
+
+ if color[0] != '#' {
+ return false, "color must start with '#'"
+ }
+
+ for _, c := range color {
+ if !strings.ContainsRune(hex, c) {
+ return false, "color must start with '#' and contain exactly 6 digits 0-9, a-f, A-F"
+ }
+ }
+
+ return true, ""
+}
diff --git a/internal/server/api/models.go b/internal/server/api/models.go
new file mode 100644
index 0000000..317ee6e
--- /dev/null
+++ b/internal/server/api/models.go
@@ -0,0 +1,15 @@
+package api
+
+const (
+ MaxIconSize = 16
+ DefaultFrequencyName = "main"
+)
+
+const (
+ PermNoAccess = 0 + iota
+ PermRead
+ PermReadWrite
+ PermMax
+)
+
+
diff --git a/internal/server/server.go b/internal/server/server.go
index 5f4376a..c15414f 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -289,6 +289,8 @@ func processRequest(ctx context.Context, sess *session.Session, request packet.P
// 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.CreateNetwork:
+ return timeout(20*time.Millisecond, api.CreateNetwork, ctx, sess, request)
case *packet.SendMessage:
return timeout(20*time.Millisecond, api.SendMessage, ctx, sess, request)
case *packet.RequestMessages:
@@ -303,6 +305,7 @@ func timeout[T packet.Payload](
apiRequest func(context.Context, *session.Session, T) packet.Payload,
ctx context.Context, sess *session.Session, request T,
) packet.Payload {
+ // TODO: Remove the channel and just wait directly?
responseChan := make(chan packet.Payload)
ctx, cancel := context.WithTimeout(ctx, timeoutDuration)
defer cancel()