diff --git a/firewall/privacy_mapper.go b/firewall/privacy_mapper.go index fd077c91..91848cb0 100644 --- a/firewall/privacy_mapper.go +++ b/firewall/privacy_mapper.go @@ -336,8 +336,8 @@ func handleGetInfoResponse(db firewalldb.PrivacyMapDB, tx firewalldb.PrivacyMapTx) error { var err error - pseudoPubKey, err = firewalldb.HideString( - tx, r.IdentityPubkey, + pseudoPubKey, err = firewalldb.HideString( //nolint:lll + ctx, tx, r.IdentityPubkey, ) return err @@ -397,14 +397,14 @@ func handleFwdHistoryResponse(db firewalldb.PrivacyMapDB, if !flags.Contains(session.ClearChanIDs) { // Deterministically hide channel ids. chanIn, err = firewalldb.HideUint64( - tx, chanIn, + ctx, tx, chanIn, ) if err != nil { return err } chanOut, err = firewalldb.HideUint64( - tx, chanOut, + ctx, tx, chanOut, ) if err != nil { return err @@ -500,7 +500,7 @@ func handleFeeReportResponse(db firewalldb.PrivacyMapDB, chanID := c.ChanId if !flags.Contains(session.ClearChanIDs) { chanID, err = firewalldb.HideUint64( - tx, chanID, + ctx, tx, chanID, ) if err != nil { return err @@ -510,7 +510,7 @@ func handleFeeReportResponse(db firewalldb.PrivacyMapDB, chanPoint := c.ChannelPoint if !flags.Contains(session.ClearChanIDs) { chanPoint, err = firewalldb.HideChanPointStr( - tx, chanPoint, + ctx, tx, chanPoint, ) if err != nil { return err @@ -599,7 +599,7 @@ func handleListChannelsResponse(db firewalldb.PrivacyMapDB, remotePub := c.RemotePubkey if hidePubkeys { remotePub, err = firewalldb.HideString( - tx, c.RemotePubkey, + ctx, tx, c.RemotePubkey, ) if err != nil { return err @@ -610,14 +610,14 @@ func handleListChannelsResponse(db firewalldb.PrivacyMapDB, chanID := c.ChanId if hideChanIds { chanPoint, err = firewalldb.HideChanPointStr( - tx, c.ChannelPoint, + ctx, tx, c.ChannelPoint, ) if err != nil { return err } chanID, err = firewalldb.HideUint64( - tx, c.ChanId, + ctx, tx, c.ChanId, ) if err != nil { return err @@ -830,7 +830,7 @@ func handleUpdatePolicyResponse(db firewalldb.PrivacyMapDB, } txid, index, err := firewalldb.HideChanPoint( - tx, u.Outpoint.TxidStr, + ctx, tx, u.Outpoint.TxidStr, u.Outpoint.OutputIndex, ) if err != nil { @@ -957,7 +957,7 @@ func handleClosedChannelsResponse(db firewalldb.PrivacyMapDB, remotePub := c.RemotePubkey if !flags.Contains(session.ClearPubkeys) { remotePub, err = firewalldb.HideString( - tx, remotePub, + ctx, tx, remotePub, ) if err != nil { return err @@ -985,7 +985,7 @@ func handleClosedChannelsResponse(db firewalldb.PrivacyMapDB, channelPoint := c.ChannelPoint if !flags.Contains(session.ClearChanIDs) { channelPoint, err = firewalldb.HideChanPointStr( - tx, c.ChannelPoint, + ctx, tx, c.ChannelPoint, ) if err != nil { return err @@ -995,7 +995,7 @@ func handleClosedChannelsResponse(db firewalldb.PrivacyMapDB, chanID := c.ChanId if !flags.Contains(session.ClearChanIDs) { chanID, err = firewalldb.HideUint64( - tx, c.ChanId, + ctx, tx, c.ChanId, ) if err != nil { return err @@ -1005,7 +1005,7 @@ func handleClosedChannelsResponse(db firewalldb.PrivacyMapDB, closingTxid := c.ClosingTxHash if !flags.Contains(session.ClearClosingTxIds) { closingTxid, err = firewalldb.HideString( - tx, c.ClosingTxHash, + ctx, tx, c.ClosingTxHash, ) if err != nil { return err @@ -1052,7 +1052,8 @@ func handleClosedChannelsResponse(db firewalldb.PrivacyMapDB, // obfuscatePendingChannel is a helper to obfuscate the fields of a pending // channel. -func obfuscatePendingChannel(c *lnrpc.PendingChannelsResponse_PendingChannel, +func obfuscatePendingChannel(ctx context.Context, + c *lnrpc.PendingChannelsResponse_PendingChannel, tx firewalldb.PrivacyMapTx, randIntn func(int) (int, error), flags session.PrivacyFlags) ( *lnrpc.PendingChannelsResponse_PendingChannel, error) { @@ -1062,7 +1063,7 @@ func obfuscatePendingChannel(c *lnrpc.PendingChannelsResponse_PendingChannel, remotePub := c.RemoteNodePub if !flags.Contains(session.ClearPubkeys) { remotePub, err = firewalldb.HideString( - tx, remotePub, + ctx, tx, remotePub, ) if err != nil { return nil, err @@ -1099,7 +1100,7 @@ func obfuscatePendingChannel(c *lnrpc.PendingChannelsResponse_PendingChannel, chanPoint := c.ChannelPoint if !flags.Contains(session.ClearChanIDs) { chanPoint, err = firewalldb.HideChanPointStr( - tx, c.ChannelPoint, + ctx, tx, c.ChannelPoint, ) if err != nil { return nil, err @@ -1163,7 +1164,7 @@ func handlePendingChannelsResponse(db firewalldb.PrivacyMapDB, var err error pendingChannel, err := obfuscatePendingChannel( - c.Channel, tx, randIntn, flags, + ctx, c.Channel, tx, randIntn, flags, ) if err != nil { return err @@ -1187,7 +1188,7 @@ func handlePendingChannelsResponse(db firewalldb.PrivacyMapDB, var err error pendingChannel, err := obfuscatePendingChannel( - c.Channel, tx, randIntn, flags, + ctx, c.Channel, tx, randIntn, flags, ) if err != nil { return err @@ -1195,8 +1196,8 @@ func handlePendingChannelsResponse(db firewalldb.PrivacyMapDB, closingTxid := c.ClosingTxid if !flags.Contains(session.ClearClosingTxIds) { - closingTxid, err = firewalldb.HideString( - tx, c.ClosingTxid, + closingTxid, err = firewalldb.HideString( //nolint:lll + ctx, tx, c.ClosingTxid, ) if err != nil { return err @@ -1216,7 +1217,7 @@ func handlePendingChannelsResponse(db firewalldb.PrivacyMapDB, var err error pendingChannel, err := obfuscatePendingChannel( - c.Channel, tx, randIntn, flags, + ctx, c.Channel, tx, randIntn, flags, ) if err != nil { return err @@ -1225,7 +1226,7 @@ func handlePendingChannelsResponse(db firewalldb.PrivacyMapDB, closingTxid := c.ClosingTxid if !flags.Contains(session.ClearClosingTxIds) { closingTxid, err = firewalldb.HideString( - tx, c.ClosingTxid, + ctx, tx, c.ClosingTxid, ) if err != nil { return err @@ -1277,7 +1278,7 @@ func handlePendingChannelsResponse(db firewalldb.PrivacyMapDB, var err error pendingChannel, err := obfuscatePendingChannel( - c.Channel, tx, randIntn, flags, + ctx, c.Channel, tx, randIntn, flags, ) if err != nil { return err @@ -1297,7 +1298,7 @@ func handlePendingChannelsResponse(db firewalldb.PrivacyMapDB, closingTxid := c.ClosingTxid if !flags.Contains(session.ClearClosingTxIds) { closingTxid, err = firewalldb.HideString( - tx, closingTxid, + ctx, tx, closingTxid, ) if err != nil { return err @@ -1314,7 +1315,7 @@ func handlePendingChannelsResponse(db firewalldb.PrivacyMapDB, ) { closingTxHex, err = firewalldb.HideString( - tx, closingTxHex, + ctx, tx, closingTxHex, ) if err != nil { return err @@ -1454,8 +1455,9 @@ func handleBatchOpenChannelResponse(db firewalldb.PrivacyMapDB, return err } - txID, outIdx, err := firewalldb.HideChanPoint( - tx, txId.String(), p.OutputIndex, + txID, outIdx, err := firewalldb.HideChanPoint( //nolint:lll + ctx, tx, txId.String(), + p.OutputIndex, ) if err != nil { return err @@ -1600,7 +1602,7 @@ func handleChannelOpenResponse(db firewalldb.PrivacyMapDB, if !flags.Contains(session.ClearChanIDs) { txid, index, err = firewalldb.HideChanPoint( - tx, txid, index, + ctx, tx, txid, index, ) if err != nil { return err diff --git a/firewall/privacy_mapper_test.go b/firewall/privacy_mapper_test.go index 24582f8c..7c84054b 100644 --- a/firewall/privacy_mapper_test.go +++ b/firewall/privacy_mapper_test.go @@ -1077,7 +1077,7 @@ func newMockDB(t *testing.T, preloadRealToPseudo map[string]string, tx firewalldb.PrivacyMapTx) error { for r, p := range preloadRealToPseudo { - require.NoError(t, tx.NewPair(r, p)) + require.NoError(t, tx.NewPair(ctx, r, p)) } return nil }) @@ -1121,7 +1121,9 @@ func (m *mockPrivacyMapDB) View(ctx context.Context, return f(ctx, m) } -func (m *mockPrivacyMapDB) NewPair(real, pseudo string) error { +func (m *mockPrivacyMapDB) NewPair(_ context.Context, real, + pseudo string) error { + m.r2p[real] = pseudo m.p2r[pseudo] = real return nil diff --git a/firewalldb/privacy_mapper.go b/firewalldb/privacy_mapper.go index e4fee472..fb9524b4 100644 --- a/firewalldb/privacy_mapper.go +++ b/firewalldb/privacy_mapper.go @@ -71,7 +71,7 @@ type PrivacyMapDB interface { // real-pseudo pairs. type PrivacyMapTx interface { // NewPair persists a new real-pseudo pair. - NewPair(real, pseudo string) error + NewPair(ctx context.Context, real, pseudo string) error // PseudoToReal returns the real value associated with the given pseudo // value. If no such pair is found, then ErrNoSuchKeyFound is returned. @@ -181,7 +181,7 @@ type privacyMapTx struct { // NewPair inserts a new real-pseudo pair into the db. // // NOTE: this is part of the PrivacyMapTx interface. -func (p *privacyMapTx) NewPair(real, pseudo string) error { +func (p *privacyMapTx) NewPair(_ context.Context, real, pseudo string) error { privacyBucket, err := getBucket(p.boltTx, privacyBucketKey) if err != nil { return err @@ -314,7 +314,9 @@ func (p *privacyMapTx) FetchAllPairs() (*PrivacyMapPairs, error) { return NewPrivacyMapPairs(pairs), nil } -func HideString(tx PrivacyMapTx, real string) (string, error) { +func HideString(ctx context.Context, tx PrivacyMapTx, real string) (string, + error) { + pseudo, err := tx.RealToPseudo(real) if err != nil && err != ErrNoSuchKeyFound { return "", err @@ -328,7 +330,7 @@ func HideString(tx PrivacyMapTx, real string) (string, error) { return "", err } - if err = tx.NewPair(real, pseudo); err != nil { + if err = tx.NewPair(ctx, real, pseudo); err != nil { return "", err } @@ -360,7 +362,9 @@ func RevealString(tx PrivacyMapTx, pseudo string) (string, error) { return tx.PseudoToReal(pseudo) } -func HideUint64(tx PrivacyMapTx, real uint64) (uint64, error) { +func HideUint64(ctx context.Context, tx PrivacyMapTx, real uint64) (uint64, + error) { + str := Uint64ToStr(real) pseudo, err := tx.RealToPseudo(str) if err != nil && err != ErrNoSuchKeyFound { @@ -371,7 +375,7 @@ func HideUint64(tx PrivacyMapTx, real uint64) (uint64, error) { } pseudoUint64, pseudoUint64Str := NewPseudoUint64() - if err := tx.NewPair(str, pseudoUint64Str); err != nil { + if err := tx.NewPair(ctx, str, pseudoUint64Str); err != nil { return 0, err } @@ -391,8 +395,8 @@ func RevealUint64(tx PrivacyMapTx, pseudo uint64) (uint64, error) { return StrToUint64(real) } -func HideChanPoint(tx PrivacyMapTx, txid string, index uint32) (string, - uint32, error) { +func HideChanPoint(ctx context.Context, tx PrivacyMapTx, txid string, + index uint32) (string, uint32, error) { cp := fmt.Sprintf("%s:%d", txid, index) pseudo, err := tx.RealToPseudo(cp) @@ -408,7 +412,7 @@ func HideChanPoint(tx PrivacyMapTx, txid string, index uint32) (string, return "", 0, err } - if err := tx.NewPair(cp, newCp); err != nil { + if err := tx.NewPair(ctx, cp, newCp); err != nil { return "", 0, err } @@ -444,13 +448,15 @@ func NewPseudoUint32() uint32 { return binary.BigEndian.Uint32(b) } -func HideChanPointStr(tx PrivacyMapTx, cp string) (string, error) { +func HideChanPointStr(ctx context.Context, tx PrivacyMapTx, cp string) (string, + error) { + txid, index, err := DecodeChannelPoint(cp) if err != nil { return "", err } - newTxid, newIndex, err := HideChanPoint(tx, txid, index) + newTxid, newIndex, err := HideChanPoint(ctx, tx, txid, index) if err != nil { return "", err } @@ -458,10 +464,12 @@ func HideChanPointStr(tx PrivacyMapTx, cp string) (string, error) { return fmt.Sprintf("%s:%d", newTxid, newIndex), nil } -func HideBytes(tx PrivacyMapTx, realBytes []byte) ([]byte, error) { +func HideBytes(ctx context.Context, tx PrivacyMapTx, realBytes []byte) ([]byte, + error) { + real := hex.EncodeToString(realBytes) - pseudo, err := HideString(tx, real) + pseudo, err := HideString(ctx, tx, real) if err != nil { return nil, err } diff --git a/firewalldb/privacy_mapper_test.go b/firewalldb/privacy_mapper_test.go index 7a48881a..03a8584e 100644 --- a/firewalldb/privacy_mapper_test.go +++ b/firewalldb/privacy_mapper_test.go @@ -29,7 +29,7 @@ func TestPrivacyMapStorage(t *testing.T) { _, err = tx.PseudoToReal("pseudo") require.ErrorIs(t, err, ErrNoSuchKeyFound) - err = tx.NewPair("real", "pseudo") + err = tx.NewPair(ctx, "real", "pseudo") require.NoError(t, err) pseudo, err := tx.RealToPseudo("real") @@ -59,7 +59,7 @@ func TestPrivacyMapStorage(t *testing.T) { _, err = tx.PseudoToReal("pseudo") require.ErrorIs(t, err, ErrNoSuchKeyFound) - err = tx.NewPair("real 2", "pseudo 2") + err = tx.NewPair(ctx, "real 2", "pseudo 2") require.NoError(t, err) pseudo, err := tx.RealToPseudo("real 2") @@ -90,29 +90,29 @@ func TestPrivacyMapStorage(t *testing.T) { require.Empty(t, m.pairs) // Add a new pair. - err = tx.NewPair("real 1", "pseudo 1") + err = tx.NewPair(ctx, "real 1", "pseudo 1") require.NoError(t, err) // Try to add a new pair that has the same real value as the // first pair. This should fail. - err = tx.NewPair("real 1", "pseudo 2") + err = tx.NewPair(ctx, "real 1", "pseudo 2") require.ErrorContains(t, err, "an entry already exists for "+ "real value") // Try to add a new pair that has the same pseudo value as the // first pair. This should fail. - err = tx.NewPair("real 2", "pseudo 1") + err = tx.NewPair(ctx, "real 2", "pseudo 1") require.ErrorContains(t, err, "an entry already exists for "+ "pseudo value") // Add a few more pairs. - err = tx.NewPair("real 2", "pseudo 2") + err = tx.NewPair(ctx, "real 2", "pseudo 2") require.NoError(t, err) - err = tx.NewPair("real 3", "pseudo 3") + err = tx.NewPair(ctx, "real 3", "pseudo 3") require.NoError(t, err) - err = tx.NewPair("real 4", "pseudo 4") + err = tx.NewPair(ctx, "real 4", "pseudo 4") require.NoError(t, err) // Check that FetchAllPairs correctly returns all the pairs. @@ -201,7 +201,7 @@ func TestPrivacyMapTxs(t *testing.T) { err = pdb1.Update(ctx, func(ctx context.Context, tx PrivacyMapTx) error { - err := tx.NewPair("real", "pseudo") + err := tx.NewPair(ctx, "real", "pseudo") if err != nil { return err } diff --git a/session_rpcserver.go b/session_rpcserver.go index 43055f65..139a14b8 100644 --- a/session_rpcserver.go +++ b/session_rpcserver.go @@ -1226,11 +1226,11 @@ func (s *sessionRpcServer) AddAutopilotSession(ctx context.Context, // Register all the privacy map pairs for this session ID. privDB := s.cfg.privMap(sess.GroupID) - err = privDB.Update(ctx, func(_ context.Context, + err = privDB.Update(ctx, func(ctx context.Context, tx firewalldb.PrivacyMapTx) error { for r, p := range newPrivMapPairs { - err := tx.NewPair(r, p) + err := tx.NewPair(ctx, r, p) if err != nil { return err }