summaryrefslogtreecommitdiff
path: root/internal/server
diff options
context:
space:
mode:
Diffstat (limited to 'internal/server')
-rw-r--r--internal/server/api/api.go97
-rw-r--r--internal/server/server.go21
2 files changed, 114 insertions, 4 deletions
diff --git a/internal/server/api/api.go b/internal/server/api/api.go
index bddc43d..2fa013a 100644
--- a/internal/server/api/api.go
+++ b/internal/server/api/api.go
@@ -1016,3 +1016,100 @@ func EditMessage(ctx context.Context, sess *session.Session, request *packet.Edi
assert.Never("unreachable")
return nil
}
+
+func TrustUser(ctx context.Context, sess *session.Session, request *packet.TrustUser) packet.Payload {
+ if sess.ID() == request.User {
+ return &packet.Error{Error: "you cannot trust yourself"}
+ }
+
+ queries := data.New(db)
+
+ user, err := queries.GetUserById(ctx, request.User)
+ if err == sql.ErrNoRows {
+ return &packet.Error{Error: "requested user doesn't exist"}
+ }
+ if err != nil {
+ log.Println("database error 0:", err)
+ return &ErrInternalError
+ }
+
+ if request.Trust {
+ publicKey, err := queries.GetTrustedPublicKey(ctx, data.GetTrustedPublicKeyParams{
+ TrustingUserID: sess.ID(),
+ TrustedUserID: user.ID,
+ })
+ if err == nil {
+ return &packet.TrustInfo{
+ Trusteds: []snowflake.ID{user.ID},
+ TrustedPublicKeys: []ed25519.PublicKey{publicKey},
+ RemovedTrusteds: nil,
+ }
+ }
+ if err != nil && err != sql.ErrNoRows {
+ log.Println("database error 1:", err)
+ return &ErrInternalError
+ }
+
+ err = queries.TrustUser(ctx, data.TrustUserParams{
+ TrustingUserID: sess.ID(),
+ TrustedUserID: user.ID,
+ TrustedPublicKey: user.PublicKey,
+ })
+ if err != nil && err != sql.ErrNoRows {
+ log.Println("database error 2:", err)
+ return &ErrInternalError
+ }
+
+ return &packet.TrustInfo{
+ Trusteds: []snowflake.ID{user.ID},
+ TrustedPublicKeys: []ed25519.PublicKey{user.PublicKey},
+ RemovedTrusteds: nil,
+ }
+ } else {
+ err = queries.UntrustUser(ctx, data.UntrustUserParams{
+ TrustingUserID: sess.ID(),
+ TrustedUserID: user.ID,
+ })
+ if err == sql.ErrNoRows {
+ return &packet.TrustInfo{
+ Trusteds: nil,
+ TrustedPublicKeys: nil,
+ RemovedTrusteds: []snowflake.ID{user.ID},
+ }
+ }
+ if err != nil {
+ log.Println("database error 3:", err)
+ return &ErrInternalError
+ }
+
+ return &packet.TrustInfo{
+ Trusteds: nil,
+ TrustedPublicKeys: nil,
+ RemovedTrusteds: []snowflake.ID{user.ID},
+ }
+ }
+}
+
+func GetUserTrusteds(ctx context.Context, sess *session.Session) packet.Payload {
+ queries := data.New(db)
+
+ trustedRows, err := queries.GetUserTrusteds(ctx, sess.ID())
+ if err != nil && err != sql.ErrNoRows {
+ log.Println("database error 0:", err)
+ return &ErrInternalError
+ }
+
+ trusteds := make([]snowflake.ID, 0, len(trustedRows))
+ trustedPublicKeys := make([]ed25519.PublicKey, 0, len(trustedRows))
+
+ for _, row := range trustedRows {
+ trusteds = append(trusteds, row.TrustedUserID)
+ trustedPublicKeys = append(trustedPublicKeys, row.TrustedPublicKey)
+ }
+
+ return &packet.TrustInfo{
+ Trusteds: trusteds,
+ TrustedPublicKeys: trustedPublicKeys,
+ RemovedTrusteds: nil,
+ }
+}
diff --git a/internal/server/server.go b/internal/server/server.go
index 8e69d70..3dcc16b 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -242,15 +242,25 @@ func (server *server) handleConnection(conn net.Conn) {
// Send initial packets
payload := api.GetUserData(ctx, sess, &packet.GetUserData{})
- dataPacket := packet.NewPacket(packet.NewMsgPackEncoder(payload))
- sess.Write(ctx, dataPacket)
+ if payload == &api.ErrInternalError {
+ return // closes the connection
+ }
+ pkt := packet.NewPacket(packet.NewMsgPackEncoder(payload))
+ sess.Write(ctx, pkt)
+
+ payload = api.GetUserTrusteds(ctx, sess)
+ if payload == &api.ErrInternalError {
+ return // closes the connection
+ }
+ pkt = packet.NewPacket(packet.NewMsgPackEncoder(payload))
+ sess.Write(ctx, pkt)
payload, err = api.GetNetworksInfo(ctx, sess)
if err != nil {
return // closes the connection
}
- infoPacket := packet.NewPacket(packet.NewMsgPackEncoder(payload))
- sess.Write(ctx, infoPacket)
+ pkt = packet.NewPacket(packet.NewMsgPackEncoder(payload))
+ sess.Write(ctx, pkt)
// Infinite read loop
buffer := make([]byte, 512)
@@ -371,6 +381,9 @@ func processRequest(ctx context.Context, sess *session.Session, request packet.P
case *packet.SetMember:
response = timeout(50*time.Millisecond, api.SetMember, ctx, sess, request)
+ case *packet.TrustUser:
+ response = timeout(10*time.Millisecond, api.TrustUser, ctx, sess, request)
+
default:
response = &packet.Error{Error: "use of disallowed packet type for request"}
}