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}