summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--internal/data/users.sql.go37
-rw-r--r--internal/packet/types.go9
-rw-r--r--internal/server/api/api.go46
-rw-r--r--internal/server/server.go8
-rw-r--r--query/users.sql5
5 files changed, 96 insertions, 9 deletions
diff --git a/internal/data/users.sql.go b/internal/data/users.sql.go
index 33eab3f..562ab35 100644
--- a/internal/data/users.sql.go
+++ b/internal/data/users.sql.go
@@ -109,6 +109,43 @@ func (q *Queries) GetUserByPublicKey(ctx context.Context, publicKey ed25519.Publ
return i, err
}
+const getUserNetworks = `-- name: GetUserNetworks :many
+SELECT networks.id, networks.owner_id, networks.name, networks.icon, networks.bg_hex_color, networks.fg_hex_color, networks.is_public FROM networks
+JOIN users_networks ON networks.id = users_networks.network_id
+WHERE users_networks.user_id = ?
+`
+
+func (q *Queries) GetUserNetworks(ctx context.Context, userID snowflake.ID) ([]Network, error) {
+ rows, err := q.db.QueryContext(ctx, getUserNetworks, userID)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ var items []Network
+ for rows.Next() {
+ var i Network
+ if err := rows.Scan(
+ &i.ID,
+ &i.OwnerID,
+ &i.Name,
+ &i.Icon,
+ &i.BgHexColor,
+ &i.FgHexColor,
+ &i.IsPublic,
+ ); err != nil {
+ return nil, err
+ }
+ items = append(items, i)
+ }
+ if err := rows.Close(); err != nil {
+ return nil, err
+ }
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return items, nil
+}
+
const setUserDescription = `-- name: SetUserDescription :one
UPDATE users SET
description = ?
diff --git a/internal/packet/types.go b/internal/packet/types.go
index 98db72a..897d4e7 100644
--- a/internal/packet/types.go
+++ b/internal/packet/types.go
@@ -65,17 +65,10 @@ func (m *SetNetworkUser) Type() PacketType {
return PacketSetNetworkUser
}
-type Member struct {
- JoinedAt string
- User data.User
- IsAdmin bool
- IsMuted bool
-}
-
type FullNetwork struct {
data.Network
Frequencies []data.Frequency
- Members []Member
+ Members []data.GetNetworkMembersRow
}
type NetworksInfo struct {
diff --git a/internal/server/api/api.go b/internal/server/api/api.go
index 50176b6..c67f906 100644
--- a/internal/server/api/api.go
+++ b/internal/server/api/api.go
@@ -102,6 +102,7 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C
log.Println("database error:", err)
return &internalError
}
+ defer tx.Rollback() //nolint
queries := data.New(db)
qtx := queries.WithTx(tx)
@@ -161,7 +162,7 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C
fullNetwork := packet.FullNetwork{
Network: network,
Frequencies: []data.Frequency{frequency},
- Members: []packet.Member{{
+ Members: []data.GetNetworkMembersRow{{
JoinedAt: networkUser.JoinedAt,
User: user,
IsAdmin: networkUser.IsAdmin,
@@ -172,3 +173,46 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C
Networks: []packet.FullNetwork{fullNetwork},
}
}
+
+func GetNetworksInfo(ctx context.Context, sess *session.Session) (packet.Payload, error) {
+ networksInfo := &packet.NetworksInfo{}
+
+ tx, err := db.BeginTx(ctx, nil)
+ if err != nil {
+ return nil, err
+ }
+ defer tx.Rollback() //nolint
+
+ queries := data.New(db)
+ qtx := queries.WithTx(tx)
+
+ networks, err := qtx.GetUserNetworks(ctx, sess.ID())
+ if err != nil {
+ return nil, err
+ }
+
+ for _, network := range networks {
+ frequencies, err := qtx.GetNetworkFrequencies(ctx, network.ID)
+ if err != nil {
+ return nil, err
+ }
+
+ members, err := qtx.GetNetworkMembers(ctx, network.ID)
+ if err != nil {
+ return nil, err
+ }
+
+ networksInfo.Networks = append(networksInfo.Networks, packet.FullNetwork{
+ Network: network,
+ Frequencies: frequencies,
+ Members: members,
+ })
+ }
+
+ err = tx.Commit()
+ if err != nil {
+ return nil, err
+ }
+
+ return networksInfo, nil
+}
diff --git a/internal/server/server.go b/internal/server/server.go
index db3cb6b..474be10 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -221,6 +221,14 @@ func (server *server) handleConnection(conn net.Conn) {
}
}()
+ // Send initial packets
+ payload, err := api.GetNetworksInfo(server.ctx, sess)
+ if err != nil {
+ return // closes the connection
+ }
+ infoPacket := packet.NewPacket(packet.NewMsgPackEncoder(payload))
+ sess.Write(server.ctx, infoPacket)
+
buffer := make([]byte, 512)
for {
n, err := conn.Read(buffer)
diff --git a/query/users.sql b/query/users.sql
index d136f15..2528513 100644
--- a/query/users.sql
+++ b/query/users.sql
@@ -46,3 +46,8 @@ RETURNING *;
UPDATE users SET
is_deleted = true
WHERE id = ? AND is_deleted = false;
+
+-- name: GetUserNetworks :many
+SELECT networks.* FROM networks
+JOIN users_networks ON networks.id = users_networks.network_id
+WHERE users_networks.user_id = ?;