diff options
Diffstat (limited to 'internal/server')
| -rw-r--r-- | internal/server/api/api.go | 93 | ||||
| -rw-r--r-- | internal/server/api/helpers.go | 23 | ||||
| -rw-r--r-- | internal/server/api/models.go | 15 | ||||
| -rw-r--r-- | internal/server/server.go | 3 |
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() |
