firewall: map session ID to group ID

This commit is contained in:
Elle Mouton 2023-06-19 15:56:16 +02:00
parent 921874608c
commit 3c837a9bc6
No known key found for this signature in database
GPG key ID: D7D916376026F177
2 changed files with 31 additions and 13 deletions

View file

@ -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{

View file

@ -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,