Skip to main content

allwright/
client_hook.rs

1use std::marker::PhantomData;
2use std::sync::Arc;
3use std::sync::atomic::{AtomicU64, Ordering};
4use std::{
5    fs,
6    io::{Read, Write},
7};
8
9use crate::proto::context_session_command::Command as ContextCommand;
10use crate::proto::context_session_event::Event as ContextEvent;
11use crate::proto::hook_completed_event::Result as HookCompletionResult;
12use crate::proto::register_hook_command::Hook as RegisterHook;
13use crate::proto::{
14    ContextSessionCommand, ReadFileChunkCommand, RegisterDownloadHook, RegisterFileChooserHook,
15    RegisterHookCommand, RegisterNewPageHook, SaveDownloadCommand, SetFileChooserFilesCommand,
16    UploadFileChunkCommand, WaitForHookCommand,
17};
18use std::path::Path;
19use tokio::sync::Mutex as AsyncMutex;
20
21use super::command::command_retry_options;
22use super::tab::ensure_tab_open;
23use super::types::{CommandOptions, Error, Result, Tab, TabInner, TabState};
24
25static TRANSFER_COUNTER: AtomicU64 = AtomicU64::new(1);
26
27pub trait HookType: private::Sealed + Clone + Send + Sync + 'static {
28    type Output;
29}
30
31mod private {
32    use super::*;
33
34    pub trait Sealed {
35        fn name() -> &'static str;
36        fn decode(page: &Tab, result: HookCompletionResult) -> Result<<Self as HookType>::Output>
37        where
38            Self: HookType;
39    }
40}
41
42#[derive(Debug, Clone, Copy, Default)]
43pub struct NewPage;
44
45pub const NEW_PAGE: NewPage = NewPage;
46
47impl HookType for NewPage {
48    type Output = Tab;
49}
50
51impl private::Sealed for NewPage {
52    fn name() -> &'static str {
53        "new_page"
54    }
55
56    fn decode(page: &Tab, result: HookCompletionResult) -> Result<<Self as HookType>::Output> {
57        let HookCompletionResult::NewPage(new_page) = result else {
58            return Err(Error::new("new page hook completed with an invalid result"));
59        };
60        Ok(Tab {
61            inner: Arc::new(TabInner {
62                runtime: Arc::clone(&page.inner.runtime),
63                surface_session_id: page.inner.surface_session_id.clone(),
64                session_id: new_page.context_session_id,
65                state: AsyncMutex::new(TabState::default()),
66            }),
67        })
68    }
69}
70
71#[derive(Debug, Clone, Copy, Default)]
72pub struct FileChooserHook;
73
74pub const FILE_CHOOSER: FileChooserHook = FileChooserHook;
75
76impl HookType for FileChooserHook {
77    type Output = FileChooser;
78}
79
80impl private::Sealed for FileChooserHook {
81    fn name() -> &'static str {
82        "file_chooser"
83    }
84
85    fn decode(page: &Tab, result: HookCompletionResult) -> Result<<Self as HookType>::Output> {
86        let HookCompletionResult::FileChooser(file_chooser) = result else {
87            return Err(Error::new(
88                "file chooser hook completed with an invalid result",
89            ));
90        };
91        Ok(FileChooser {
92            page: page.clone(),
93            id: file_chooser.file_chooser_id,
94            is_multiple: file_chooser.is_multiple,
95        })
96    }
97}
98
99#[derive(Clone)]
100pub struct FileChooser {
101    page: Tab,
102    id: String,
103    is_multiple: bool,
104}
105
106impl FileChooser {
107    pub fn id(&self) -> &str {
108        &self.id
109    }
110
111    pub fn is_multiple(&self) -> bool {
112        self.is_multiple
113    }
114
115    pub fn page(&self) -> &Tab {
116        &self.page
117    }
118
119    pub async fn set_file(&self, file: impl AsRef<Path>) -> Result<()> {
120        self.set_files([file]).await
121    }
122
123    pub async fn set_files<I, P>(&self, files: I) -> Result<()>
124    where
125        I: IntoIterator<Item = P>,
126        P: AsRef<Path>,
127    {
128        let files = files
129            .into_iter()
130            .map(|file| file.as_ref().to_path_buf())
131            .collect::<Vec<_>>();
132        let mut state = self.page.inner.state.lock().await;
133        let handle = self.page.ensure_handle(&mut state).await?;
134        ensure_tab_open(handle, &self.page.inner.session_id)?;
135        let mut file_ids = Vec::with_capacity(files.len());
136        for path in files {
137            let name = path
138                .file_name()
139                .and_then(|name| name.to_str())
140                .ok_or_else(|| Error::new("upload path requires a valid file name"))?
141                .to_string();
142            let transfer_id = format!(
143                "rust-upload-{}-{}",
144                std::process::id(),
145                TRANSFER_COUNTER.fetch_add(1, Ordering::Relaxed)
146            );
147            let mut source = fs::File::open(&path)
148                .map_err(|error| Error::new(format!("open upload {}: {error}", path.display())))?;
149            let size = source
150                .metadata()
151                .map_err(|error| Error::new(format!("inspect upload {}: {error}", path.display())))?
152                .len();
153            let mut offset = 0_u64;
154            loop {
155                let mut data = vec![0; 256 * 1024];
156                let count = source.read(&mut data).map_err(|error| {
157                    Error::new(format!("read upload {}: {error}", path.display()))
158                })?;
159                data.truncate(count);
160                let last = offset + count as u64 >= size;
161                handle
162                    .command_tx
163                    .send(ContextSessionCommand {
164                        surface_session_id: self.page.inner.surface_session_id.clone(),
165                        context_session_id: self.page.inner.session_id.clone(),
166                        command: Some(ContextCommand::UploadFileChunk(UploadFileChunkCommand {
167                            transfer_id: transfer_id.clone(),
168                            name: name.clone(),
169                            offset,
170                            data,
171                            last,
172                        })),
173                    })
174                    .await
175                    .map_err(|_| Error::new("failed to send UploadFileChunkCommand"))?;
176                offset += count as u64;
177                if last {
178                    break;
179                }
180            }
181            loop {
182                let event = handle.events.message().await?.ok_or_else(|| {
183                    Error::new("page session closed while uploading chooser file")
184                })?;
185                match event.event {
186                    Some(ContextEvent::FileUploaded(result))
187                        if result.transfer_id == transfer_id =>
188                    {
189                        file_ids.push(result.file_id);
190                        break;
191                    }
192                    Some(ContextEvent::Error(error)) => {
193                        return Err(Error::new(format!(
194                            "page session error while uploading chooser file: {}",
195                            error.message
196                        )));
197                    }
198                    _ => {}
199                }
200            }
201        }
202        handle
203            .command_tx
204            .send(ContextSessionCommand {
205                surface_session_id: self.page.inner.surface_session_id.clone(),
206                context_session_id: self.page.inner.session_id.clone(),
207                command: Some(ContextCommand::SetFileChooserFiles(
208                    SetFileChooserFilesCommand {
209                        file_chooser_id: self.id.clone(),
210                        file_ids,
211                        retry_options: None,
212                    },
213                )),
214            })
215            .await
216            .map_err(|_| Error::new("failed to send SetFileChooserFilesCommand"))?;
217        loop {
218            let event = handle
219                .events
220                .message()
221                .await?
222                .ok_or_else(|| Error::new("page session closed while setting chooser files"))?;
223            match event.event {
224                Some(ContextEvent::FileChooserFilesSet(result))
225                    if result.file_chooser_id == self.id =>
226                {
227                    return Ok(());
228                }
229                Some(ContextEvent::Error(error)) => {
230                    return Err(Error::new(format!(
231                        "page session error while setting chooser files: {}",
232                        error.message
233                    )));
234                }
235                _ => {}
236            }
237        }
238    }
239}
240
241#[derive(Debug, Clone, Copy, Default)]
242pub struct DownloadHook;
243
244pub const DOWNLOAD: DownloadHook = DownloadHook;
245
246impl HookType for DownloadHook {
247    type Output = Download;
248}
249
250impl private::Sealed for DownloadHook {
251    fn name() -> &'static str {
252        "download"
253    }
254
255    fn decode(page: &Tab, result: HookCompletionResult) -> Result<<Self as HookType>::Output> {
256        let HookCompletionResult::Download(download) = result else {
257            return Err(Error::new("download hook completed with an invalid result"));
258        };
259        Ok(Download {
260            page: page.clone(),
261            id: download.download_id,
262            url: download.url,
263            suggested_filename: download.suggested_filename,
264        })
265    }
266}
267
268#[derive(Clone)]
269pub struct Download {
270    page: Tab,
271    id: String,
272    url: String,
273    suggested_filename: String,
274}
275
276impl Download {
277    pub fn id(&self) -> &str {
278        &self.id
279    }
280
281    pub fn page(&self) -> &Tab {
282        &self.page
283    }
284
285    pub fn url(&self) -> &str {
286        &self.url
287    }
288
289    pub fn suggested_filename(&self) -> &str {
290        &self.suggested_filename
291    }
292
293    pub async fn save_as(&self, path: impl AsRef<Path>) -> Result<()> {
294        self.save_as_with_options(path, CommandOptions::default())
295            .await
296    }
297
298    pub async fn save_as_with_options(
299        &self,
300        path: impl AsRef<Path>,
301        options: CommandOptions,
302    ) -> Result<()> {
303        let mut state = self.page.inner.state.lock().await;
304        let handle = self.page.ensure_handle(&mut state).await?;
305        ensure_tab_open(handle, &self.page.inner.session_id)?;
306        handle
307            .command_tx
308            .send(ContextSessionCommand {
309                surface_session_id: self.page.inner.surface_session_id.clone(),
310                context_session_id: self.page.inner.session_id.clone(),
311                command: Some(ContextCommand::SaveDownload(SaveDownloadCommand {
312                    download_id: self.id.clone(),
313                    retry_options: command_retry_options(options.timeout_ms),
314                })),
315            })
316            .await
317            .map_err(|_| Error::new("failed to send SaveDownloadCommand"))?;
318        loop {
319            let event = handle
320                .events
321                .message()
322                .await?
323                .ok_or_else(|| Error::new("page session closed while saving download"))?;
324            match event.event {
325                Some(ContextEvent::DownloadSaved(result)) if result.download_id == self.id => {
326                    let destination = path.as_ref();
327                    let temporary = destination.with_extension(format!(
328                        "allwright-{}.tmp",
329                        TRANSFER_COUNTER.fetch_add(1, Ordering::Relaxed)
330                    ));
331                    let mut output = fs::OpenOptions::new()
332                        .write(true)
333                        .create_new(true)
334                        .open(&temporary)
335                        .map_err(|error| Error::new(format!("create download file: {error}")))?;
336                    let mut offset = 0_u64;
337                    loop {
338                        handle
339                            .command_tx
340                            .send(ContextSessionCommand {
341                                surface_session_id: self.page.inner.surface_session_id.clone(),
342                                context_session_id: self.page.inner.session_id.clone(),
343                                command: Some(ContextCommand::ReadFileChunk(
344                                    ReadFileChunkCommand {
345                                        file_id: result.file_id.clone(),
346                                        offset,
347                                        max_bytes: 256 * 1024,
348                                    },
349                                )),
350                            })
351                            .await
352                            .map_err(|_| Error::new("failed to send ReadFileChunkCommand"))?;
353                        let chunk = handle.events.message().await?.ok_or_else(|| {
354                            Error::new("page session closed while downloading file")
355                        })?;
356                        match chunk.event {
357                            Some(ContextEvent::FileChunk(chunk))
358                                if chunk.file_id == result.file_id && chunk.offset == offset =>
359                            {
360                                output.write_all(&chunk.data).map_err(|error| {
361                                    Error::new(format!("write download file: {error}"))
362                                })?;
363                                offset += chunk.data.len() as u64;
364                                if chunk.last {
365                                    drop(output);
366                                    fs::rename(&temporary, destination).map_err(|error| {
367                                        Error::new(format!("finish download file: {error}"))
368                                    })?;
369                                    return Ok(());
370                                }
371                            }
372                            Some(ContextEvent::Error(error)) => {
373                                let _ = fs::remove_file(&temporary);
374                                return Err(Error::new(format!(
375                                    "page session error while downloading file: {}",
376                                    error.message
377                                )));
378                            }
379                            _ => {}
380                        }
381                    }
382                }
383                Some(ContextEvent::Error(error)) => {
384                    return Err(Error::new(format!(
385                        "page session error while saving download: {}",
386                        error.message
387                    )));
388                }
389                _ => {}
390            }
391        }
392    }
393}
394
395pub struct Hook<T: HookType> {
396    page: Tab,
397    id: String,
398    _type: PhantomData<T>,
399}
400
401impl<T: HookType> Hook<T> {
402    pub fn id(&self) -> &str {
403        &self.id
404    }
405
406    pub async fn wait(&self) -> Result<T::Output> {
407        self.wait_with_options(CommandOptions::default()).await
408    }
409
410    pub async fn wait_with_options(&self, options: CommandOptions) -> Result<T::Output> {
411        let mut state = self.page.inner.state.lock().await;
412        let handle = self.page.ensure_handle(&mut state).await?;
413        ensure_tab_open(handle, &self.page.inner.session_id)?;
414        handle
415            .command_tx
416            .send(ContextSessionCommand {
417                surface_session_id: self.page.inner.surface_session_id.clone(),
418                context_session_id: self.page.inner.session_id.clone(),
419                command: Some(ContextCommand::WaitForHook(WaitForHookCommand {
420                    hook_id: self.id.clone(),
421                    retry_options: command_retry_options(options.timeout_ms),
422                })),
423            })
424            .await
425            .map_err(|_| Error::new("failed to send WaitForHookCommand"))?;
426
427        loop {
428            let event = handle
429                .events
430                .message()
431                .await?
432                .ok_or_else(|| Error::new("page session closed while waiting for hook"))?;
433            match event.event {
434                Some(ContextEvent::HookCompleted(completed)) if completed.hook_id == self.id => {
435                    let result = completed
436                        .result
437                        .ok_or_else(|| Error::new("hook completed without a result"))?;
438                    return T::decode(&self.page, result);
439                }
440                Some(ContextEvent::Error(error)) => {
441                    return Err(Error::new(format!(
442                        "page session error while waiting for hook: {}",
443                        error.message
444                    )));
445                }
446                _ => {}
447            }
448        }
449    }
450}
451
452impl Tab {
453    pub async fn register_hook<T: HookType>(&self, _hook_type: T) -> Result<Hook<T>> {
454        let mut state = self.inner.state.lock().await;
455        let handle = self.ensure_handle(&mut state).await?;
456        ensure_tab_open(handle, &self.inner.session_id)?;
457        handle
458            .command_tx
459            .send(ContextSessionCommand {
460                surface_session_id: self.inner.surface_session_id.clone(),
461                context_session_id: self.inner.session_id.clone(),
462                command: Some(ContextCommand::RegisterHook(RegisterHookCommand {
463                    hook: match T::name() {
464                        "new_page" => Some(RegisterHook::NewPage(RegisterNewPageHook {})),
465                        "file_chooser" => {
466                            Some(RegisterHook::FileChooser(RegisterFileChooserHook {}))
467                        }
468                        "download" => Some(RegisterHook::Download(RegisterDownloadHook {})),
469                        _ => return Err(Error::new("unsupported hook type")),
470                    },
471                })),
472            })
473            .await
474            .map_err(|_| Error::new("failed to send RegisterHookCommand"))?;
475
476        loop {
477            let event = handle
478                .events
479                .message()
480                .await?
481                .ok_or_else(|| Error::new("page session closed while registering hook"))?;
482            match event.event {
483                Some(ContextEvent::HookRegistered(registered)) => {
484                    return Ok(Hook {
485                        page: self.clone(),
486                        id: registered.hook_id,
487                        _type: PhantomData,
488                    });
489                }
490                Some(ContextEvent::Error(error)) => {
491                    return Err(Error::new(format!(
492                        "page session error while registering hook: {}",
493                        error.message
494                    )));
495                }
496                _ => {}
497            }
498        }
499    }
500}