diff --git a/firewall/rule_enforcer.go b/firewall/rule_enforcer.go index c51d2f9f..db304b7a 100644 --- a/firewall/rule_enforcer.go +++ b/firewall/rule_enforcer.go @@ -30,6 +30,7 @@ var _ mid.RequestInterceptor = (*RuleEnforcer)(nil) type RuleEnforcer struct { ruleDB firewalldb.RulesDB actionsDB firewalldb.ActionReadDBGetter + sessionIDIndexDB session.IDToGroupIndex markActionErrored func(reqID uint64, reason string) error newPrivMap firewalldb.NewPrivacyMapDB @@ -50,8 +51,9 @@ type featurePerms func(ctx context.Context) (map[string]map[string]bool, error) // NewRuleEnforcer constructs a new RuleEnforcer instance. func NewRuleEnforcer(ruleDB firewalldb.RulesDB, - actionsDB firewalldb.ActionReadDBGetter, getFeaturePerms featurePerms, - permsMgr *perms.Manager, nodeID [33]byte, + actionsDB firewalldb.ActionReadDBGetter, + sessionIDIndex session.IDToGroupIndex, + getFeaturePerms featurePerms, permsMgr *perms.Manager, nodeID [33]byte, routerClient lndclient.RouterClient, lndClient lndclient.LightningClient, ruleMgrs rules.ManagerSet, markActionErrored func(reqID uint64, reason string) error, @@ -68,6 +70,7 @@ func NewRuleEnforcer(ruleDB firewalldb.RulesDB, ruleMgrs: ruleMgrs, markActionErrored: markActionErrored, newPrivMap: privMap, + sessionIDIndexDB: sessionIDIndex, } } @@ -221,7 +224,12 @@ func (r *RuleEnforcer) handleRequest(ctx context.Context, return nil, fmt.Errorf("could not extract ID from macaroon") } - rules, err := r.collectEnforcers(ri, sessionID) + groupID, err := r.sessionIDIndexDB.GetGroupID(sessionID) + if err != nil { + return nil, err + } + + rules, err := r.collectEnforcers(ri, groupID) if err != nil { return nil, fmt.Errorf("error parsing rules: %v", err) } @@ -261,7 +269,12 @@ func (r *RuleEnforcer) handleResponse(ctx context.Context, return nil, fmt.Errorf("could not extract ID from macaroon") } - enforcers, err := r.collectEnforcers(ri, sessionID) + groupID, err := r.sessionIDIndexDB.GetGroupID(sessionID) + if err != nil { + return nil, err + } + + enforcers, err := r.collectEnforcers(ri, groupID) if err != nil { return nil, fmt.Errorf("error parsing rules: %v", err) } @@ -295,7 +308,12 @@ func (r *RuleEnforcer) handleErrorResponse(ctx context.Context, return nil, fmt.Errorf("could not extract ID from macaroon") } - enforcers, err := r.collectEnforcers(ri, sessionID) + groupID, err := r.sessionIDIndexDB.GetGroupID(sessionID) + if err != nil { + return nil, err + } + + enforcers, err := r.collectEnforcers(ri, groupID) if err != nil { return nil, fmt.Errorf("error parsing rules: %v", err) } @@ -320,7 +338,7 @@ func (r *RuleEnforcer) handleErrorResponse(ctx context.Context, // collectRule initialises and returns all the Rules that need to be enforced // for the given request. -func (r *RuleEnforcer) collectEnforcers(ri *RequestInfo, sessionID session.ID) ( +func (r *RuleEnforcer) collectEnforcers(ri *RequestInfo, groupID session.ID) ( []rules.Enforcer, error) { ruleEnforcers := make( @@ -331,7 +349,7 @@ func (r *RuleEnforcer) collectEnforcers(ri *RequestInfo, sessionID session.ID) ( for rule, value := range ri.Rules.FeatureRules[ri.MetaInfo.Feature] { r, err := r.initRule( ri.RequestID, rule, []byte(value), ri.MetaInfo.Feature, - sessionID, false, ri.WithPrivacy, + groupID, false, ri.WithPrivacy, ) if err != nil { return nil, err @@ -345,7 +363,7 @@ func (r *RuleEnforcer) collectEnforcers(ri *RequestInfo, sessionID session.ID) ( // initRule initialises a rule.Rule with any required config values. func (r *RuleEnforcer) initRule(reqID uint64, name string, value []byte, - featureName string, sessionID session.ID, sessionRule, + featureName string, groupID session.ID, sessionRule, privacy bool) (rules.Enforcer, error) { ruleValues, err := r.ruleMgrs.InitRuleValues(name, value) @@ -354,7 +372,7 @@ func (r *RuleEnforcer) initRule(reqID uint64, name string, value []byte, } if privacy { - privMap := r.newPrivMap(sessionID) + privMap := r.newPrivMap(groupID) ruleValues, err = ruleValues.PseudoToReal(privMap) if err != nil { return nil, fmt.Errorf("could not prepare rule "+ @@ -362,13 +380,13 @@ func (r *RuleEnforcer) initRule(reqID uint64, name string, value []byte, } } - allActionsDB := r.actionsDB.GetActionsReadDB(sessionID, featureName) + allActionsDB := r.actionsDB.GetActionsReadDB(groupID, featureName) actionsDB := allActionsDB.GroupFeatureActionsDB() - rulesDB := r.ruleDB.GetKVStores(name, sessionID, featureName) + rulesDB := r.ruleDB.GetKVStores(name, groupID, featureName) if sessionRule { actionsDB = allActionsDB.GroupActionsDB() - rulesDB = r.ruleDB.GetKVStores(name, sessionID, "") + rulesDB = r.ruleDB.GetKVStores(name, groupID, "") } cfg := &rules.ConfigImpl{ diff --git a/terminal.go b/terminal.go index d92a134b..8125ca1d 100644 --- a/terminal.go +++ b/terminal.go @@ -824,7 +824,7 @@ func (g *LightningTerminal) startInternalSubServers( if !g.cfg.Autopilot.Disable { ruleEnforcer := firewall.NewRuleEnforcer( - g.firewallDB, g.firewallDB, + g.firewallDB, g.firewallDB, g.sessionDB, g.autopilotClient.ListFeaturePerms, g.permsMgr, g.lndClient.NodePubkey, g.lndClient.Router,