Skip to main content

shared/extensions/
manager.rs

1use crate::{
2    State,
3    extensions::{
4        ConstructedExtension, ExtensionPermissionsBuilder, ExtensionRouteBuilder,
5        commands::CliCommandGroupBuilder,
6    },
7    settings::ExtensionPermissions,
8};
9use std::{
10    collections::{BTreeMap, HashSet},
11    sync::Arc,
12};
13use tokio::sync::{RwLock, RwLockReadGuard};
14
15pub struct ExtensionManager {
16    vec: RwLock<Vec<ConstructedExtension>>,
17    disabled: parking_lot::RwLock<Vec<compact_str::CompactString>>,
18}
19
20impl ExtensionManager {
21    pub fn new(vec: Vec<ConstructedExtension>) -> Self {
22        Self {
23            vec: RwLock::new(vec),
24            disabled: parking_lot::RwLock::new(Vec::new()),
25        }
26    }
27
28    /// Sets the extensions whose entrypoints are skipped, this has to happen before `init`
29    /// and before extension migrations are collected.
30    pub fn set_disabled(&self, disabled: Vec<compact_str::CompactString>) {
31        *self.disabled.write() = disabled;
32    }
33
34    #[inline]
35    pub fn disabled(&self) -> Vec<compact_str::CompactString> {
36        self.disabled.read().clone()
37    }
38
39    #[inline]
40    pub fn is_disabled(&self, package_name: &str) -> bool {
41        self.disabled.read().iter().any(|d| d == package_name)
42    }
43
44    pub async fn init(
45        &self,
46        state: State,
47    ) -> (
48        ExtensionRouteBuilder,
49        super::background_tasks::BackgroundTaskBuilder,
50        super::shutdown_handlers::ShutdownHandlerBuilder,
51    ) {
52        let mut route_builder = ExtensionRouteBuilder::new(state.clone());
53        let mut email_templates_builder =
54            super::email_templates::ExtensionEmailTemplateBuilder::default();
55        let mut background_tasks_builder =
56            super::background_tasks::BackgroundTaskBuilder::new(state.clone());
57        let mut shutdown_handlers_builder =
58            super::shutdown_handlers::ShutdownHandlerBuilder::new(state.clone());
59        let mut permissions_builder = ExtensionPermissionsBuilder::new(
60            crate::permissions::BASE_USER_PERMISSIONS.clone(),
61            crate::permissions::BASE_ADMIN_PERMISSIONS.clone(),
62            crate::permissions::BASE_SERVER_PERMISSIONS.clone(),
63        );
64
65        for ext in self.vec.read().await.iter() {
66            if self.is_disabled(ext.package_name) {
67                tracing::info!(extension = %ext.package_name, "extension is disabled, skipping its entrypoints");
68                continue;
69            }
70
71            let deserializer = ext.settings_deserializer(state.clone()).await;
72            crate::settings::SETTINGS_DESER_EXTENSIONS
73                .write()
74                .insert(ext.package_name, deserializer);
75        }
76        state.settings.invalidate_cache().await;
77
78        let mut contributed_permissions = BTreeMap::new();
79
80        for ext in self.vec.write().await.iter_mut() {
81            if self.is_disabled(ext.package_name) {
82                continue;
83            }
84
85            let package_name = ext.package_name;
86            let ext = match Arc::get_mut(&mut ext.extension) {
87                Some(ext) => ext,
88                None => {
89                    panic!(
90                        "Failed to get mutable reference to extension {package_name}. This should NEVER happen."
91                    );
92                }
93            };
94
95            ext.initialize(state.clone()).await;
96
97            route_builder = ext.initialize_router(state.clone(), route_builder).await;
98            email_templates_builder = ext
99                .initialize_email_templates(state.clone(), email_templates_builder)
100                .await;
101            background_tasks_builder = ext
102                .initialize_background_tasks(state.clone(), background_tasks_builder)
103                .await;
104            shutdown_handlers_builder = ext
105                .initialize_shutdown_handlers(state.clone(), shutdown_handlers_builder)
106                .await;
107
108            let before = permissions_builder.snapshot();
109            permissions_builder = ext
110                .initialize_permissions(state.clone(), permissions_builder)
111                .await;
112
113            contributed_permissions.insert(
114                compact_str::CompactString::from(package_name),
115                permissions_builder.contributions_since(&before),
116            );
117        }
118
119        *state.mail.templates.templates.write() = email_templates_builder.finish();
120
121        crate::permissions::USER_PERMISSIONS
122            .write()
123            .replace(permissions_builder.user_permissions);
124        crate::permissions::ADMIN_PERMISSIONS
125            .write()
126            .replace(permissions_builder.admin_permissions);
127        crate::permissions::SERVER_PERMISSIONS
128            .write()
129            .replace(permissions_builder.server_permissions);
130
131        self.apply_permission_snapshots(&state, contributed_permissions)
132            .await;
133
134        (
135            route_builder,
136            background_tasks_builder,
137            shutdown_handlers_builder,
138        )
139    }
140
141    /// Keeps the permissions of installed extensions valid across a disable/enable cycle:
142    /// what an enabled extension contributed is remembered, what a disabled one contributed
143    /// is fed back as inert so roles, subusers and api keys holding it still validate.
144    async fn apply_permission_snapshots(
145        &self,
146        state: &State,
147        contributed: BTreeMap<compact_str::CompactString, ExtensionPermissions>,
148    ) {
149        let stored = match state.settings.get().await {
150            Ok(settings) => settings.extension_permissions.clone(),
151            Err(err) => {
152                tracing::error!("failed to read stored extension permissions: {:#?}", err);
153                return;
154            }
155        };
156
157        let mut snapshots = BTreeMap::new();
158        let mut inert_user = HashSet::new();
159        let mut inert_admin = HashSet::new();
160        let mut inert_server = HashSet::new();
161
162        for ext in self.vec.read().await.iter() {
163            let package_name = compact_str::CompactString::from(ext.package_name);
164
165            let permissions = match contributed.get(&package_name) {
166                Some(permissions) => permissions.clone(),
167                None => match stored.get(&package_name) {
168                    Some(permissions) => {
169                        inert_user.extend(permissions.user.iter().map(ToString::to_string));
170                        inert_admin.extend(permissions.admin.iter().map(ToString::to_string));
171                        inert_server.extend(permissions.server.iter().map(ToString::to_string));
172
173                        permissions.clone()
174                    }
175                    None => continue,
176                },
177            };
178
179            snapshots.insert(package_name, permissions);
180        }
181
182        tracing::debug!(
183            count = inert_user.len() + inert_admin.len() + inert_server.len(),
184            "keeping permissions of disabled extensions grantable"
185        );
186
187        crate::permissions::USER_PERMISSIONS
188            .write()
189            .set_inert(inert_user);
190        crate::permissions::ADMIN_PERMISSIONS
191            .write()
192            .set_inert(inert_admin);
193        crate::permissions::SERVER_PERMISSIONS
194            .write()
195            .set_inert(inert_server);
196
197        if snapshots != stored
198            && let Err(err) = state.settings.set_extension_permissions(&snapshots).await
199        {
200            tracing::error!("failed to store extension permissions: {:#?}", err);
201        }
202    }
203
204    pub async fn init_cli(
205        &self,
206        env: Option<&Arc<crate::env::Env>>,
207        mut builder: CliCommandGroupBuilder,
208    ) -> CliCommandGroupBuilder {
209        for ext in self.vec.write().await.iter_mut() {
210            let ext = match Arc::get_mut(&mut ext.extension) {
211                Some(ext) => ext,
212                None => {
213                    panic!(
214                        "Failed to get mutable reference to extension {}. This should NEVER happen.",
215                        ext.package_name
216                    );
217                }
218            };
219
220            builder = ext.initialize_cli(env, builder).await;
221        }
222
223        builder
224    }
225
226    #[inline]
227    pub async fn extensions(&self) -> RwLockReadGuard<'_, Vec<ConstructedExtension>> {
228        self.vec.read().await
229    }
230
231    #[inline]
232    pub fn blocking_extensions(&self) -> RwLockReadGuard<'_, Vec<ConstructedExtension>> {
233        self.vec.blocking_read()
234    }
235
236    pub async fn call(
237        &self,
238        name: impl AsRef<str>,
239        args: &[super::ExtensionCallValue],
240    ) -> Option<super::ExtensionCallValue> {
241        for ext in self.extensions().await.iter() {
242            if self.is_disabled(ext.package_name) {
243                continue;
244            }
245
246            if let Some(ret) = ext.process_call(name.as_ref(), args).await {
247                return Some(ret);
248            }
249        }
250
251        None
252    }
253}