diff options
Diffstat (limited to 'internal/server')
| -rw-r--r-- | internal/server/api/api.go | 97 | ||||
| -rw-r--r-- | internal/server/server.go | 21 |
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"} } |
