Skip to main content

allwright/
client_hook.rs

1use std::marker::PhantomData;
2use std::sync::Arc;
3
4use crate::proto::context_session_command::Command as ContextCommand;
5use crate::proto::context_session_event::Event as ContextEvent;
6use crate::proto::hook_completed_event::Result as HookCompletionResult;
7use crate::proto::register_hook_command::Hook as RegisterHook;
8use crate::proto::{
9    ContextSessionCommand, RegisterDownloadHook, RegisterFileChooserHook, RegisterHookCommand,
10    RegisterNewPageHook, SaveDownloadCommand, SetFileChooserFilesCommand, WaitForHookCommand,
11};
12use std::path::Path;
13use tokio::sync::Mutex as AsyncMutex;
14
15use super::command::command_retry_options;
16use super::tab::ensure_tab_open;
17use super::types::{CommandOptions, Error, Result, Tab, TabInner, TabState};
18
19pub trait HookType: private::Sealed + Clone + Send + Sync + 'static {
20    type Output;
21}
22
23mod private {
24    use super::*;
25
26    pub trait Sealed {
27        fn name() -> &'static str;
28        fn decode(page: &Tab, result: HookCompletionResult) -> Result<<Self as HookType>::Output>
29        where
30            Self: HookType;
31    }
32}
33
34#[derive(Debug, Clone, Copy, Default)]
35pub struct NewPage;
36
37pub const NEW_PAGE: NewPage = NewPage;
38
39impl HookType for NewPage {
40    type Output = Tab;
41}
42
43impl private::Sealed for NewPage {
44    fn name() -> &'static str {
45        "new_page"
46    }
47
48    fn decode(page: &Tab, result: HookCompletionResult) -> Result<<Self as HookType>::Output> {
49        let HookCompletionResult::NewPage(new_page) = result else {
50            return Err(Error::new("new page hook completed with an invalid result"));
51        };
52        Ok(Tab {
53            inner: Arc::new(TabInner {
54                runtime: Arc::clone(&page.inner.runtime),
55                surface_session_id: page.inner.surface_session_id.clone(),
56                session_id: new_page.context_session_id,
57                state: AsyncMutex::new(TabState::default()),
58            }),
59        })
60    }
61}
62
63#[derive(Debug, Clone, Copy, Default)]
64pub struct FileChooserHook;
65
66pub const FILE_CHOOSER: FileChooserHook = FileChooserHook;
67
68impl HookType for FileChooserHook {
69    type Output = FileChooser;
70}
71
72impl private::Sealed for FileChooserHook {
73    fn name() -> &'static str {
74        "file_chooser"
75    }
76
77    fn decode(page: &Tab, result: HookCompletionResult) -> Result<<Self as HookType>::Output> {
78        let HookCompletionResult::FileChooser(file_chooser) = result else {
79            return Err(Error::new(
80                "file chooser hook completed with an invalid result",
81            ));
82        };
83        Ok(FileChooser {
84            page: page.clone(),
85            id: file_chooser.file_chooser_id,
86            is_multiple: file_chooser.is_multiple,
87        })
88    }
89}
90
91#[derive(Clone)]
92pub struct FileChooser {
93    page: Tab,
94    id: String,
95    is_multiple: bool,
96}
97
98impl FileChooser {
99    pub fn id(&self) -> &str {
100        &self.id
101    }
102
103    pub fn is_multiple(&self) -> bool {
104        self.is_multiple
105    }
106
107    pub fn page(&self) -> &Tab {
108        &self.page
109    }
110
111    pub async fn set_file(&self, file: impl AsRef<Path>) -> Result<()> {
112        self.set_files([file]).await
113    }
114
115    pub async fn set_files<I, P>(&self, files: I) -> Result<()>
116    where
117        I: IntoIterator<Item = P>,
118        P: AsRef<Path>,
119    {
120        let files = files
121            .into_iter()
122            .map(|file| file.as_ref().to_string_lossy().to_string())
123            .collect::<Vec<_>>();
124        let mut state = self.page.inner.state.lock().await;
125        let handle = self.page.ensure_handle(&mut state).await?;
126        ensure_tab_open(handle, &self.page.inner.session_id)?;
127        handle
128            .command_tx
129            .send(ContextSessionCommand {
130                surface_session_id: self.page.inner.surface_session_id.clone(),
131                context_session_id: self.page.inner.session_id.clone(),
132                command: Some(ContextCommand::SetFileChooserFiles(
133                    SetFileChooserFilesCommand {
134                        file_chooser_id: self.id.clone(),
135                        files,
136                        retry_options: None,
137                    },
138                )),
139            })
140            .await
141            .map_err(|_| Error::new("failed to send SetFileChooserFilesCommand"))?;
142        loop {
143            let event = handle
144                .events
145                .message()
146                .await?
147                .ok_or_else(|| Error::new("page session closed while setting chooser files"))?;
148            match event.event {
149                Some(ContextEvent::FileChooserFilesSet(result))
150                    if result.file_chooser_id == self.id =>
151                {
152                    return Ok(());
153                }
154                Some(ContextEvent::Error(error)) => {
155                    return Err(Error::new(format!(
156                        "page session error while setting chooser files: {}",
157                        error.message
158                    )));
159                }
160                _ => {}
161            }
162        }
163    }
164}
165
166#[derive(Debug, Clone, Copy, Default)]
167pub struct DownloadHook;
168
169pub const DOWNLOAD: DownloadHook = DownloadHook;
170
171impl HookType for DownloadHook {
172    type Output = Download;
173}
174
175impl private::Sealed for DownloadHook {
176    fn name() -> &'static str {
177        "download"
178    }
179
180    fn decode(page: &Tab, result: HookCompletionResult) -> Result<<Self as HookType>::Output> {
181        let HookCompletionResult::Download(download) = result else {
182            return Err(Error::new("download hook completed with an invalid result"));
183        };
184        Ok(Download {
185            page: page.clone(),
186            id: download.download_id,
187            url: download.url,
188            suggested_filename: download.suggested_filename,
189        })
190    }
191}
192
193#[derive(Clone)]
194pub struct Download {
195    page: Tab,
196    id: String,
197    url: String,
198    suggested_filename: String,
199}
200
201impl Download {
202    pub fn id(&self) -> &str {
203        &self.id
204    }
205
206    pub fn page(&self) -> &Tab {
207        &self.page
208    }
209
210    pub fn url(&self) -> &str {
211        &self.url
212    }
213
214    pub fn suggested_filename(&self) -> &str {
215        &self.suggested_filename
216    }
217
218    pub async fn save_as(&self, path: impl AsRef<Path>) -> Result<()> {
219        self.save_as_with_options(path, CommandOptions::default())
220            .await
221    }
222
223    pub async fn save_as_with_options(
224        &self,
225        path: impl AsRef<Path>,
226        options: CommandOptions,
227    ) -> Result<()> {
228        let mut state = self.page.inner.state.lock().await;
229        let handle = self.page.ensure_handle(&mut state).await?;
230        ensure_tab_open(handle, &self.page.inner.session_id)?;
231        handle
232            .command_tx
233            .send(ContextSessionCommand {
234                surface_session_id: self.page.inner.surface_session_id.clone(),
235                context_session_id: self.page.inner.session_id.clone(),
236                command: Some(ContextCommand::SaveDownload(SaveDownloadCommand {
237                    download_id: self.id.clone(),
238                    path: path.as_ref().to_string_lossy().to_string(),
239                    retry_options: command_retry_options(options.timeout_ms),
240                })),
241            })
242            .await
243            .map_err(|_| Error::new("failed to send SaveDownloadCommand"))?;
244        loop {
245            let event = handle
246                .events
247                .message()
248                .await?
249                .ok_or_else(|| Error::new("page session closed while saving download"))?;
250            match event.event {
251                Some(ContextEvent::DownloadSaved(result)) if result.download_id == self.id => {
252                    return Ok(());
253                }
254                Some(ContextEvent::Error(error)) => {
255                    return Err(Error::new(format!(
256                        "page session error while saving download: {}",
257                        error.message
258                    )));
259                }
260                _ => {}
261            }
262        }
263    }
264}
265
266pub struct Hook<T: HookType> {
267    page: Tab,
268    id: String,
269    _type: PhantomData<T>,
270}
271
272impl<T: HookType> Hook<T> {
273    pub fn id(&self) -> &str {
274        &self.id
275    }
276
277    pub async fn wait(&self) -> Result<T::Output> {
278        self.wait_with_options(CommandOptions::default()).await
279    }
280
281    pub async fn wait_with_options(&self, options: CommandOptions) -> Result<T::Output> {
282        let mut state = self.page.inner.state.lock().await;
283        let handle = self.page.ensure_handle(&mut state).await?;
284        ensure_tab_open(handle, &self.page.inner.session_id)?;
285        handle
286            .command_tx
287            .send(ContextSessionCommand {
288                surface_session_id: self.page.inner.surface_session_id.clone(),
289                context_session_id: self.page.inner.session_id.clone(),
290                command: Some(ContextCommand::WaitForHook(WaitForHookCommand {
291                    hook_id: self.id.clone(),
292                    retry_options: command_retry_options(options.timeout_ms),
293                })),
294            })
295            .await
296            .map_err(|_| Error::new("failed to send WaitForHookCommand"))?;
297
298        loop {
299            let event = handle
300                .events
301                .message()
302                .await?
303                .ok_or_else(|| Error::new("page session closed while waiting for hook"))?;
304            match event.event {
305                Some(ContextEvent::HookCompleted(completed)) if completed.hook_id == self.id => {
306                    let result = completed
307                        .result
308                        .ok_or_else(|| Error::new("hook completed without a result"))?;
309                    return T::decode(&self.page, result);
310                }
311                Some(ContextEvent::Error(error)) => {
312                    return Err(Error::new(format!(
313                        "page session error while waiting for hook: {}",
314                        error.message
315                    )));
316                }
317                _ => {}
318            }
319        }
320    }
321}
322
323impl Tab {
324    pub async fn register_hook<T: HookType>(&self, _hook_type: T) -> Result<Hook<T>> {
325        let mut state = self.inner.state.lock().await;
326        let handle = self.ensure_handle(&mut state).await?;
327        ensure_tab_open(handle, &self.inner.session_id)?;
328        handle
329            .command_tx
330            .send(ContextSessionCommand {
331                surface_session_id: self.inner.surface_session_id.clone(),
332                context_session_id: self.inner.session_id.clone(),
333                command: Some(ContextCommand::RegisterHook(RegisterHookCommand {
334                    hook: match T::name() {
335                        "new_page" => Some(RegisterHook::NewPage(RegisterNewPageHook {})),
336                        "file_chooser" => {
337                            Some(RegisterHook::FileChooser(RegisterFileChooserHook {}))
338                        }
339                        "download" => Some(RegisterHook::Download(RegisterDownloadHook {})),
340                        _ => return Err(Error::new("unsupported hook type")),
341                    },
342                })),
343            })
344            .await
345            .map_err(|_| Error::new("failed to send RegisterHookCommand"))?;
346
347        loop {
348            let event = handle
349                .events
350                .message()
351                .await?
352                .ok_or_else(|| Error::new("page session closed while registering hook"))?;
353            match event.event {
354                Some(ContextEvent::HookRegistered(registered)) => {
355                    return Ok(Hook {
356                        page: self.clone(),
357                        id: registered.hook_id,
358                        _type: PhantomData,
359                    });
360                }
361                Some(ContextEvent::Error(error)) => {
362                    return Err(Error::new(format!(
363                        "page session error while registering hook: {}",
364                        error.message
365                    )));
366                }
367                _ => {}
368            }
369        }
370    }
371}