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}