diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/packet/packet.go | 125 | ||||
| -rw-r--r-- | internal/packet/types.go | 46 | ||||
| -rw-r--r-- | internal/server/api/api.go | 20 | ||||
| -rw-r--r-- | internal/server/ctxkeys/ctxkeys.go | 5 | ||||
| -rw-r--r-- | internal/server/server.go | 4 |
5 files changed, 122 insertions, 78 deletions
diff --git a/internal/packet/packet.go b/internal/packet/packet.go index d3dcd7e..2285065 100644 --- a/internal/packet/packet.go +++ b/internal/packet/packet.go @@ -52,6 +52,14 @@ type PacketType uint8 const ( PacketError PacketType = iota + PacketTosInfo + PacketAcceptTos + + PacketGetNonce + PacketNonceInfo + + PacketAuthenticate + PacketSetUserData PacketGetUserData @@ -92,7 +100,57 @@ const ( PacketMax ) -func Init() { +var packetNames = map[PacketType]string{ + PacketError: "PacketError", + + PacketTosInfo: "PacketTosInfo", + PacketAcceptTos: "PacketAcceptTos", + + PacketGetNonce: "PacketGetNonce", + PacketNonceInfo: "PacketNonceInfo", + + PacketAuthenticate: "PacketAuthenticate", + + PacketSetUserData: "PacketSetUserData", + PacketGetUserData: "PacketGetUserData", + + PacketCreateNetwork: "PacketCreateNetwork", + PacketUpdateNetwork: "PacketUpdateNetwork", + PacketTransferNetwork: "PacketTransferNetwork", + PacketDeleteNetwork: "PacketDeleteNetwork", + PacketNetworksInfo: "PacketNetworksInfo", + + PacketCreateFrequency: "PacketCreateFrequency", + PacketUpdateFrequency: "PacketUpdateFrequency", + PacketDeleteFrequency: "PacketDeleteFrequency", + PacketSwapFrequencies: "PacketSwapFrequencies", + PacketFrequenciesInfo: "PacketFrequenciesInfo", + + PacketSendMessage: "PacketSendMessage", + PacketEditMessage: "PacketEditMessage", + PacketDeleteMessage: "PacketDeleteMessage", + PacketRequestMessages: "PacketRequestMessages", + PacketMessagesInfo: "PacketMessagesInfo", + + PacketGetBannedMembers: "PacketGetBannedMembers", + PacketSetMember: "PacketSetMember", + PacketMembersInfo: "PacketMembersInfo", + + PacketTrustUser: "PacketTrustUser", + PacketTrustInfo: "PacketTrustInfo", + + PacketSetLastReadMessages: "PacketSetLastReadMessages", + PacketNotificationsInfo: "PacketNotificationsInfo", + + PacketBlockUser: "PacketBlockUser", + PacketBlockInfo: "PacketBlockInfo", + + PacketGetUsers: "PacketGetUsers", + PacketUsersInfo: "PacketUsersInfo", +} + +func init() { + assert.Assert(len(packetNames) == int(PacketMax), "packetName length mismatches with PacketMax", "len(packetNames)", len(packetNames), "PacketMax", int(PacketMax)) assert.Assert(PacketMax <= 64, "packet types exceeded allowed limit of 64 types") } @@ -101,71 +159,10 @@ func (e PacketType) IsSupported() bool { } func (e PacketType) String() string { - switch e { - case PacketBlockInfo: - return "PacketBlockInfo" - case PacketBlockUser: - return "PacketBlockUser" - case PacketCreateFrequency: - return "PacketCreateFrequency" - case PacketCreateNetwork: - return "PacketCreateNetwork" - case PacketDeleteFrequency: - return "PacketDeleteFrequency" - case PacketDeleteMessage: - return "PacketDeleteMessage" - case PacketDeleteNetwork: - return "PacketDeleteNetwork" - case PacketEditMessage: - return "PacketEditMessage" - case PacketError: - return "PacketError" - case PacketFrequenciesInfo: - return "PacketFrequenciesInfo" - case PacketGetBannedMembers: - return "PacketGetBannedMembers" - case PacketGetUserData: - return "PacketGetUserData" - case PacketGetUsers: - return "PacketGetUsers" - case PacketMax: - return "PacketMax" - case PacketMembersInfo: - return "PacketMembersInfo" - case PacketMessagesInfo: - return "PacketMessagesInfo" - case PacketNetworksInfo: - return "PacketNetworksInfo" - case PacketNotificationsInfo: - return "PacketNotificationsInfo" - case PacketRequestMessages: - return "PacketRequestMessages" - case PacketSendMessage: - return "PacketSendMessage" - case PacketSetLastReadMessages: - return "PacketSetLastReadMessages" - case PacketSetMember: - return "PacketSetMember" - case PacketSetUserData: - return "PacketSetUserData" - case PacketSwapFrequencies: - return "PacketSwapFrequencies" - case PacketTransferNetwork: - return "PacketTransferNetwork" - case PacketTrustInfo: - return "PacketTrustInfo" - case PacketTrustUser: - return "PacketTrustUser" - case PacketUpdateFrequency: - return "PacketUpdateFrequency" - case PacketUpdateNetwork: - return "PacketUpdateNetwork" - case PacketUsersInfo: - return "PacketUsersInfo" - default: - assert.Assert(!e.IsSupported(), "missing string for supported packet type", "type", e) + if !e.IsSupported() { return fmt.Sprintf("UnsupportedPacket(%d)", e) } + return packetNames[e] } const ( diff --git a/internal/packet/types.go b/internal/packet/types.go index f251450..156fa5f 100644 --- a/internal/packet/types.go +++ b/internal/packet/types.go @@ -193,8 +193,9 @@ func (m *MembersInfo) Type() PacketType { } type SetUserData struct { - Data *string - User *data.User + Data *string + User *data.User + Nonce []byte // optional } func (m *SetUserData) Type() PacketType { @@ -288,3 +289,44 @@ type UsersInfo struct { func (m *UsersInfo) Type() PacketType { return PacketUsersInfo } + +type TosInfo struct { + Tos string + PrivacyPolicy string + Date string // 2025-07-05 +} + +func (m *TosInfo) Type() PacketType { + return PacketTosInfo +} + +type AcceptTos struct { + IAgreeToTheTermsOfServiceAndPrivacyPolicy bool +} + +func (m *AcceptTos) Type() PacketType { + return PacketAcceptTos +} + +type GetNonce struct{} + +func (m *GetNonce) Type() PacketType { + return PacketGetNonce +} + +type NonceInfo struct { + Nonce []byte +} + +func (m *NonceInfo) Type() PacketType { + return PacketNonceInfo +} + +type Authenticate struct { + PubKey ed25519.PublicKey + Signature []byte +} + +func (m *Authenticate) Type() PacketType { + return PacketAuthenticate +} diff --git a/internal/server/api/api.go b/internal/server/api/api.go index dddd047..068e074 100644 --- a/internal/server/api/api.go +++ b/internal/server/api/api.go @@ -8,6 +8,7 @@ import ( "errors" "fmt" "log" + "log/slog" "strconv" "strings" @@ -1343,7 +1344,7 @@ func SetLastReadMessages(ctx context.Context, sess *session.Session, request *pa tx, err := db.BeginTx(ctx, nil) if err != nil { - log.Println("database error:", err) + slog.ErrorContext(ctx, "database error", "error", err) return &ErrInternalError } defer func() { _ = tx.Rollback() }() @@ -1351,25 +1352,28 @@ func SetLastReadMessages(ctx context.Context, sess *session.Session, request *pa queries := data.New(db) qtx := queries.WithTx(tx) + // OPTIMIZE: Convert this loop into a SQL query for i := 0; i < len(request.Source); i++ { _, err := qtx.GetUserById(ctx, request.Source[i]) if err == nil { + // ID is signal err = qtx.SetLastReadMessage(ctx, data.SetLastReadMessageParams{ UserID: sess.ID(), SourceID: request.Source[i], LastRead: request.LastRead[i], }) if err != nil { - log.Println("api.go:1362 database error:", err) + slog.ErrorContext(ctx, "database error", "error", err) return &ErrInternalError } continue } - if err != nil && err != sql.ErrNoRows { - log.Println("api.go:1368 database error:", err) + if err != sql.ErrNoRows { + slog.ErrorContext(ctx, "database error", "error", err) return &ErrInternalError } + // ID is frequency frequency, err := qtx.GetFrequencyById(ctx, request.Source[i]) if err == sql.ErrNoRows { return &packet.Error{Error: fmt.Sprintf( @@ -1377,13 +1381,13 @@ func SetLastReadMessages(ctx context.Context, sess *session.Session, request *pa )} } if err != nil { - log.Println("api.go:1379 database error:", err) + slog.ErrorContext(ctx, "database error", "error", err) return &ErrInternalError } if frequency.Perms == packet.PermNoAccess { isAdmin, err := IsNetworkAdmin(ctx, qtx, sess.ID(), frequency.NetworkID) if err != nil { - log.Println("api.go:1385 database error:", err) + slog.ErrorContext(ctx, "database error", "error", err) return &ErrInternalError } if !isAdmin { @@ -1397,7 +1401,7 @@ func SetLastReadMessages(ctx context.Context, sess *session.Session, request *pa LastRead: request.LastRead[i], }) if err != nil { - log.Println("api.go:1399 database error:", err) + slog.ErrorContext(ctx, "database error", "error", err) return &ErrInternalError } continue @@ -1405,7 +1409,7 @@ func SetLastReadMessages(ctx context.Context, sess *session.Session, request *pa err = tx.Commit() if err != nil { - log.Println("api.go:1407 database error:", err) + slog.ErrorContext(ctx, "database error", "error", err) return &ErrInternalError } diff --git a/internal/server/ctxkeys/ctxkeys.go b/internal/server/ctxkeys/ctxkeys.go index 3677a07..1fb0fa6 100644 --- a/internal/server/ctxkeys/ctxkeys.go +++ b/internal/server/ctxkeys/ctxkeys.go @@ -12,9 +12,10 @@ type key int const ( UserID key = iota IpAddr - KeyMax Evicted EvictedBy + + KeyMax ) var keyNames = map[key]string{ @@ -24,7 +25,7 @@ var keyNames = map[key]string{ EvictedBy: "evicted_by", } -func Init() { +func init() { assert.Assert(len(keyNames) == int(KeyMax), "Keys in keyNames mismatch amount of keys", "len(keyNames)", len(keyNames), "KeyMax", int(KeyMax)) } diff --git a/internal/server/server.go b/internal/server/server.go index 5791f8e..1727961 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -186,11 +186,11 @@ func (server *server) handleConnection(conn net.Conn) { defer slog.InfoContext(ctx, "connection closed") defer conn.Close() - var writerWg *sync.WaitGroup + var writerWg sync.WaitGroup done := make(chan struct{}) framer := packet.NewFramer() - sess := session.NewSession(server, addr, cancel, writerWg) + sess := session.NewSession(server, addr, cancel, &writerWg) go func() { <-ctx.Done() // Remove session after cancellation |
