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, RegisterHookCommand, RegisterNewPageHook, WaitForHookCommand,
10};
11use tokio::sync::Mutex as AsyncMutex;
12
13use super::command::command_retry_options;
14use super::tab::ensure_tab_open;
15use super::types::{CommandOptions, Error, Result, Tab, TabInner, TabState};
16
17pub trait HookType: private::Sealed + Clone + Send + Sync + 'static {
18    type Output;
19}
20
21mod private {
22    use super::*;
23
24    pub trait Sealed {
25        fn name() -> &'static str;
26        fn decode(page: &Tab, result: HookCompletionResult) -> Result<<Self as HookType>::Output>
27        where
28            Self: HookType;
29    }
30}
31
32#[derive(Debug, Clone, Copy, Default)]
33pub struct NewPage;
34
35pub const NEW_PAGE: NewPage = NewPage;
36
37impl HookType for NewPage {
38    type Output = Tab;
39}
40
41impl private::Sealed for NewPage {
42    fn name() -> &'static str {
43        "new_page"
44    }
45
46    fn decode(page: &Tab, result: HookCompletionResult) -> Result<<Self as HookType>::Output> {
47        let HookCompletionResult::NewPage(new_page) = result;
48        Ok(Tab {
49            inner: Arc::new(TabInner {
50                runtime: Arc::clone(&page.inner.runtime),
51                surface_session_id: page.inner.surface_session_id.clone(),
52                session_id: new_page.context_session_id,
53                state: AsyncMutex::new(TabState::default()),
54            }),
55        })
56    }
57}
58
59pub struct Hook<T: HookType> {
60    page: Tab,
61    id: String,
62    _type: PhantomData<T>,
63}
64
65impl<T: HookType> Hook<T> {
66    pub fn id(&self) -> &str {
67        &self.id
68    }
69
70    pub async fn wait(&self) -> Result<T::Output> {
71        self.wait_with_options(CommandOptions::default()).await
72    }
73
74    pub async fn wait_with_options(&self, options: CommandOptions) -> Result<T::Output> {
75        let mut state = self.page.inner.state.lock().await;
76        let handle = self.page.ensure_handle(&mut state).await?;
77        ensure_tab_open(handle, &self.page.inner.session_id)?;
78        handle
79            .command_tx
80            .send(ContextSessionCommand {
81                surface_session_id: self.page.inner.surface_session_id.clone(),
82                context_session_id: self.page.inner.session_id.clone(),
83                command: Some(ContextCommand::WaitForHook(WaitForHookCommand {
84                    hook_id: self.id.clone(),
85                    retry_options: command_retry_options(options.timeout_ms),
86                })),
87            })
88            .await
89            .map_err(|_| Error::new("failed to send WaitForHookCommand"))?;
90
91        loop {
92            let event = handle
93                .events
94                .message()
95                .await?
96                .ok_or_else(|| Error::new("page session closed while waiting for hook"))?;
97            match event.event {
98                Some(ContextEvent::HookCompleted(completed)) if completed.hook_id == self.id => {
99                    let result = completed
100                        .result
101                        .ok_or_else(|| Error::new("hook completed without a result"))?;
102                    return T::decode(&self.page, result);
103                }
104                Some(ContextEvent::Error(error)) => {
105                    return Err(Error::new(format!(
106                        "page session error while waiting for hook: {}",
107                        error.message
108                    )));
109                }
110                _ => {}
111            }
112        }
113    }
114}
115
116impl Tab {
117    pub async fn register_hook<T: HookType>(&self, _hook_type: T) -> Result<Hook<T>> {
118        let mut state = self.inner.state.lock().await;
119        let handle = self.ensure_handle(&mut state).await?;
120        ensure_tab_open(handle, &self.inner.session_id)?;
121        handle
122            .command_tx
123            .send(ContextSessionCommand {
124                surface_session_id: self.inner.surface_session_id.clone(),
125                context_session_id: self.inner.session_id.clone(),
126                command: Some(ContextCommand::RegisterHook(RegisterHookCommand {
127                    hook: match T::name() {
128                        "new_page" => Some(RegisterHook::NewPage(RegisterNewPageHook {})),
129                        _ => return Err(Error::new("unsupported hook type")),
130                    },
131                })),
132            })
133            .await
134            .map_err(|_| Error::new("failed to send RegisterHookCommand"))?;
135
136        loop {
137            let event = handle
138                .events
139                .message()
140                .await?
141                .ok_or_else(|| Error::new("page session closed while registering hook"))?;
142            match event.event {
143                Some(ContextEvent::HookRegistered(registered)) => {
144                    return Ok(Hook {
145                        page: self.clone(),
146                        id: registered.hook_id,
147                        _type: PhantomData,
148                    });
149                }
150                Some(ContextEvent::Error(error)) => {
151                    return Err(Error::new(format!(
152                        "page session error while registering hook: {}",
153                        error.message
154                    )));
155                }
156                _ => {}
157            }
158        }
159    }
160}