systemprompt_security/authz/parent_chain/
mod.rs1mod cache;
21mod chains;
22mod sources;
23
24use std::collections::BTreeMap;
25use std::sync::Arc;
26
27use systemprompt_identifiers::{MarketplaceId, PluginId, UserId};
28
29use super::error::{AuthzError, AuthzResult};
30use super::repository::AccessControlRepository;
31use super::resolver::{ResolveInput, ResolveParent, resolve};
32use super::subject::{SubjectAttributes, SubjectDimension};
33use super::types::{AccessRule, Decision, DenyReason, EntityKind, EntityRef};
34
35pub use cache::ChainIndexCache;
36pub use sources::{ChainSources, MarketplaceSource};
37
38#[derive(Debug, Clone)]
39pub struct LoadedParent {
40 pub entity: EntityRef,
41 pub rules: Vec<AccessRule>,
42 pub default_included: Option<bool>,
43}
44
45impl LoadedParent {
46 #[must_use]
47 pub fn as_resolve_parent(&self) -> ResolveParent<'_> {
48 ResolveParent {
49 entity: &self.entity,
50 rules: &self.rules,
51 default_included: self.default_included,
52 }
53 }
54}
55
56#[derive(Debug, Clone, Copy)]
57pub struct ResolveBase<'a> {
58 pub rules: &'a [AccessRule],
59 pub user_id: &'a UserId,
60 pub user_roles: &'a [String],
61 pub default_included: Option<bool>,
62 pub attributes: &'a SubjectAttributes,
63 pub dimensions: &'a [SubjectDimension],
64}
65
66#[derive(Debug, Clone, Default)]
67pub struct ParentChainIndex {
68 pub(super) marketplaces: BTreeMap<MarketplaceId, LoadedParent>,
69 pub(super) plugins: BTreeMap<PluginId, LoadedParent>,
70 pub(super) sources: Arc<ChainSources>,
71}
72
73impl ParentChainIndex {
74 #[must_use]
75 pub const fn from_parts(
76 marketplaces: BTreeMap<MarketplaceId, LoadedParent>,
77 plugins: BTreeMap<PluginId, LoadedParent>,
78 sources: Arc<ChainSources>,
79 ) -> Self {
80 Self {
81 marketplaces,
82 plugins,
83 sources,
84 }
85 }
86
87 pub async fn load(
88 repo: &AccessControlRepository,
89 sources: Arc<ChainSources>,
90 ) -> AuthzResult<Self> {
91 let marketplaces = load_marketplaces(repo, &sources).await?;
92
93 let plugin_ids = sources.plugin_ids_to_load();
94 let mut rules = repo
95 .list_rules_bulk(EntityKind::Plugin, &plugin_ids)
96 .await?;
97 let entities = repo
98 .list_entities_bulk(EntityKind::Plugin, &plugin_ids)
99 .await?;
100 let mut plugins = BTreeMap::new();
101 for id in plugin_ids {
102 let entity = EntityRef::from_kind_and_id(EntityKind::Plugin, &id)
103 .map_err(|e| AuthzError::Validation(e.to_string()))?;
104 let parent = LoadedParent {
105 entity,
106 rules: rules.remove(&id).unwrap_or_default(),
107 default_included: entities.get(&id).map(|row| row.default_included),
108 };
109 plugins.insert(PluginId::new(id), parent);
110 }
111
112 Ok(Self {
113 marketplaces,
114 plugins,
115 sources,
116 })
117 }
118
119 #[must_use]
120 pub fn resolve(&self, kind: EntityKind, id: &str, base: ResolveBase<'_>) -> Decision {
121 let entity = match EntityRef::from_kind_and_id(kind, id) {
122 Ok(entity) => entity,
123 Err(e) => {
124 return Decision::Deny {
125 reason: DenyReason::InvalidEntity {
126 entity_kind: kind,
127 id: id.to_owned(),
128 detail: e.to_string(),
129 },
130 };
131 },
132 };
133 let resolve_with = |parents: &[ResolveParent<'_>]| {
134 resolve(ResolveInput {
135 entity: &entity,
136 rules: base.rules,
137 user_id: base.user_id,
138 user_roles: base.user_roles,
139 default_included: base.default_included,
140 parents,
141 attributes: base.attributes,
142 dimensions: base.dimensions,
143 })
144 };
145
146 let mut first_deny = None;
147 for chain in self.chains_for(kind, id) {
148 let decision = resolve_with(&chain);
149 if decision.permits() {
150 return decision;
151 }
152 first_deny.get_or_insert(decision);
153 }
154 first_deny.unwrap_or_else(|| resolve_with(&[]))
155 }
156}
157
158async fn load_marketplaces(
159 repo: &AccessControlRepository,
160 sources: &ChainSources,
161) -> AuthzResult<BTreeMap<MarketplaceId, LoadedParent>> {
162 let ids = sources.marketplace_ids_to_load();
163 if ids.is_empty() {
164 return Ok(BTreeMap::new());
165 }
166 let raw: Vec<String> = ids.iter().map(|id| id.as_str().to_owned()).collect();
167 let mut rules = repo.list_rules_bulk(EntityKind::Marketplace, &raw).await?;
168 let entities = repo
169 .list_entities_bulk(EntityKind::Marketplace, &raw)
170 .await?;
171 Ok(ids
172 .into_iter()
173 .map(|id| {
174 let fallback = sources
175 .marketplaces
176 .get(&id)
177 .and_then(|source| source.fallback_default_included);
178 let parent = LoadedParent {
179 entity: EntityRef::Marketplace(id.clone()),
180 rules: rules.remove(id.as_str()).unwrap_or_default(),
181 default_included: entities
182 .get(id.as_str())
183 .map(|row| row.default_included)
184 .or(fallback),
185 };
186 (id, parent)
187 })
188 .collect())
189}