1use std::collections::HashMap;
16use std::fmt;
17use std::marker::PhantomData;
18use std::sync::atomic::{AtomicU64, Ordering};
19use std::sync::{Arc, Mutex as StdMutex};
20
21use serde::de::DeserializeOwned;
22use serde::{Deserialize, Deserializer, Serialize, Serializer};
23use serde_json::{json, Value};
24use thiserror::Error;
25use tokio::io::{AsyncRead, AsyncWrite};
26use tokio::sync::{broadcast, oneshot, Mutex};
27use tokio::task::JoinHandle;
28
29use super::protocol::{Playwright, Root, RootInitializeParams, SDKLanguage};
30use super::transport::{read_message, write_message};
31
32pub type Binary = String;
34
35#[derive(Debug, Error)]
37pub enum ProtocolError {
38 #[error("{name}: {message}")]
40 Remote {
41 name: String,
43 message: String,
45 stack: Option<String>,
47 },
48 #[error("Playwright driver connection closed: {0}")]
50 Closed(String),
51 #[error("Playwright driver I/O error: {0}")]
53 Io(#[from] std::io::Error),
54 #[error("Playwright protocol message did not match the spec: {0}")]
56 Json(#[from] serde_json::Error),
57 #[error("unknown Playwright object {0}")]
59 UnknownObject(String),
60 #[error("Playwright object {guid} is a {actual}, expected {expected}")]
62 WrongInterface {
63 guid: String,
65 expected: &'static str,
67 actual: String,
69 },
70 #[error("Playwright driver unavailable: {0}")]
72 Driver(String),
73}
74
75impl ProtocolError {
76 pub fn is_timeout(&self) -> bool {
78 matches!(self, ProtocolError::Remote { name, .. } if name == "TimeoutError")
79 }
80}
81
82pub struct Ref<T> {
84 pub guid: String,
86 marker: PhantomData<fn() -> T>,
87}
88
89impl<T> Ref<T> {
90 pub fn new(guid: impl Into<String>) -> Self {
92 Self {
93 guid: guid.into(),
94 marker: PhantomData,
95 }
96 }
97}
98
99impl<T: ChannelType> From<&T> for Ref<T> {
100 fn from(object: &T) -> Self {
101 Ref::new(object.channel().guid())
102 }
103}
104
105impl<T> Clone for Ref<T> {
106 fn clone(&self) -> Self {
107 Ref::new(self.guid.clone())
108 }
109}
110
111impl<T> Default for Ref<T> {
112 fn default() -> Self {
113 Ref::new(String::new())
114 }
115}
116
117impl<T> PartialEq for Ref<T> {
118 fn eq(&self, other: &Self) -> bool {
119 self.guid == other.guid
120 }
121}
122
123impl<T> fmt::Debug for Ref<T> {
124 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
125 write!(f, "Ref({})", self.guid)
126 }
127}
128
129impl<T> Serialize for Ref<T> {
130 fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
131 ObjectRef {
132 guid: self.guid.clone(),
133 }
134 .serialize(serializer)
135 }
136}
137
138impl<'de, T> Deserialize<'de> for Ref<T> {
139 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
140 Ok(Ref::new(ObjectRef::deserialize(deserializer)?.guid))
141 }
142}
143
144#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
146pub struct ObjectRef {
147 pub guid: String,
149}
150
151pub trait ProtocolEvent: Sized + Send + 'static {
153 fn parse(method: &str, params: Value) -> Result<Self, serde_json::Error>;
155}
156
157pub trait ChannelType: Sized + Clone + Send + Sync + 'static {
159 const INTERFACE: &'static str;
161 type Initializer: DeserializeOwned;
163 type Event: ProtocolEvent;
165
166 fn accepts(interface: &str) -> bool;
169 fn from_channel(channel: Channel) -> Self;
171 fn channel(&self) -> &Channel;
173
174 fn guid(&self) -> &str {
176 self.channel().guid()
177 }
178
179 fn initializer(&self) -> Result<Self::Initializer, ProtocolError> {
181 let raw = self.channel().connection().raw_initializer(self.guid())?;
182 Ok(serde_json::from_value(raw)?)
183 }
184
185 fn events(&self) -> EventStream<Self::Event> {
187 EventStream {
188 guid: self.guid().to_string(),
189 receiver: self.channel().connection().subscribe(),
190 marker: PhantomData,
191 }
192 }
193}
194
195#[derive(Clone)]
197pub struct Channel {
198 guid: Arc<str>,
199 connection: Connection,
200}
201
202impl fmt::Debug for Channel {
203 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
204 write!(f, "Channel({})", self.guid)
205 }
206}
207
208impl Channel {
209 pub fn guid(&self) -> &str {
211 &self.guid
212 }
213
214 pub fn connection(&self) -> &Connection {
216 &self.connection
217 }
218
219 pub async fn send<P, R>(&self, method: &str, params: &P) -> Result<R, ProtocolError>
221 where
222 P: Serialize + ?Sized,
223 R: DeserializeOwned,
224 {
225 self.send_with_timeout(method, params, self.connection.default_timeout())
226 .await
227 }
228
229 pub async fn send_with_timeout<P, R>(
232 &self,
233 method: &str,
234 params: &P,
235 timeout_ms: Option<f64>,
236 ) -> Result<R, ProtocolError>
237 where
238 P: Serialize + ?Sized,
239 R: DeserializeOwned,
240 {
241 let params = serde_json::to_value(params)?;
242 let result = match self
243 .connection
244 .call_with_timeout(&self.guid, method, params, timeout_ms)
245 .await?
246 {
247 Value::Null => json!({}),
249 result => result,
250 };
251 Ok(serde_json::from_value(result)?)
252 }
253
254 pub async fn send_no_result<P>(&self, method: &str, params: &P) -> Result<(), ProtocolError>
256 where
257 P: Serialize + ?Sized,
258 {
259 let params = serde_json::to_value(params)?;
260 self.connection.call(&self.guid, method, params).await?;
261 Ok(())
262 }
263}
264
265#[derive(Debug, Clone, PartialEq)]
267pub struct RawEvent {
268 pub guid: String,
270 pub method: String,
272 pub params: Value,
274}
275
276pub struct EventStream<E> {
278 guid: String,
279 receiver: broadcast::Receiver<RawEvent>,
280 marker: PhantomData<fn() -> E>,
281}
282
283impl<E: ProtocolEvent> EventStream<E> {
284 pub async fn recv(&mut self) -> Result<E, ProtocolError> {
289 loop {
290 match self.receiver.recv().await {
291 Ok(event) if event.guid == self.guid => {
292 return Ok(E::parse(&event.method, event.params)?);
293 }
294 Ok(_) | Err(broadcast::error::RecvError::Lagged(_)) => continue,
295 Err(broadcast::error::RecvError::Closed) => {
296 return Err(ProtocolError::Closed("event stream ended".to_string()));
297 }
298 }
299 }
300 }
301}
302
303#[derive(Debug, Clone)]
304struct ObjectEntry {
305 interface: String,
306 initializer: Value,
307 parent: String,
308}
309
310type Pending = HashMap<u64, oneshot::Sender<Result<Value, ProtocolError>>>;
311
312struct Inner {
313 writer: Mutex<Box<dyn AsyncWrite + Send + Unpin>>,
314 next_id: AtomicU64,
315 pending: StdMutex<Pending>,
316 objects: StdMutex<HashMap<String, ObjectEntry>>,
317 events: broadcast::Sender<RawEvent>,
318 closed: StdMutex<Option<String>>,
319 default_timeout: StdMutex<Option<f64>>,
320 reader: StdMutex<Option<JoinHandle<()>>>,
321}
322
323impl Drop for Inner {
324 fn drop(&mut self) {
325 if let Some(reader) = self.reader.lock().ok().and_then(|mut slot| slot.take()) {
326 reader.abort();
327 }
328 }
329}
330
331#[derive(Clone)]
333pub struct Connection {
334 inner: Arc<Inner>,
335}
336
337impl fmt::Debug for Connection {
338 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
339 f.debug_struct("Connection")
340 .field("closed", &self.close_reason())
341 .finish()
342 }
343}
344
345const EVENT_BUFFER: usize = 1024;
346
347impl Connection {
348 pub fn new<R, W>(reader: R, writer: W) -> Self
351 where
352 R: AsyncRead + Send + Unpin + 'static,
353 W: AsyncWrite + Send + Unpin + 'static,
354 {
355 let (events, _) = broadcast::channel(EVENT_BUFFER);
356 let inner = Arc::new(Inner {
357 writer: Mutex::new(Box::new(writer)),
358 next_id: AtomicU64::new(1),
359 pending: StdMutex::new(HashMap::new()),
360 objects: StdMutex::new(HashMap::new()),
361 events,
362 closed: StdMutex::new(None),
363 default_timeout: StdMutex::new(None),
364 reader: StdMutex::new(None),
365 });
366 let weak = Arc::downgrade(&inner);
367 let task = tokio::spawn(async move {
368 let mut reader = reader;
369 let reason = loop {
370 let message = match read_message(&mut reader).await {
371 Ok(Some(message)) => message,
372 Ok(None) => break "the driver closed its output".to_string(),
373 Err(err) => break format!("reading from the driver failed: {err}"),
374 };
375 let Some(inner) = weak.upgrade() else {
376 return;
377 };
378 Connection { inner }.dispatch(message);
379 };
380 if let Some(inner) = weak.upgrade() {
381 Connection { inner }.mark_closed(reason);
382 }
383 });
384 if let Ok(mut slot) = inner.reader.lock() {
385 *slot = Some(task);
386 }
387 Self { inner }
388 }
389
390 pub async fn initialize(&self) -> Result<Playwright, ProtocolError> {
392 let root = Root::from_channel(self.channel(""));
393 let result = root
394 .initialize(RootInitializeParams {
395 sdk_language: SDKLanguage::Javascript,
396 })
397 .await?;
398 self.object(&result.playwright)
399 }
400
401 pub fn channel(&self, guid: &str) -> Channel {
403 Channel {
404 guid: Arc::from(guid),
405 connection: self.clone(),
406 }
407 }
408
409 pub fn object<T: ChannelType>(&self, reference: &Ref<T>) -> Result<T, ProtocolError> {
411 self.object_by_guid(&reference.guid)
412 }
413
414 pub fn object_by_guid<T: ChannelType>(&self, guid: &str) -> Result<T, ProtocolError> {
416 let interface = self
417 .lock_objects()
418 .get(guid)
419 .map(|entry| entry.interface.clone())
420 .ok_or_else(|| ProtocolError::UnknownObject(guid.to_string()))?;
421 if !T::accepts(&interface) {
422 return Err(ProtocolError::WrongInterface {
423 guid: guid.to_string(),
424 expected: T::INTERFACE,
425 actual: interface,
426 });
427 }
428 Ok(T::from_channel(self.channel(guid)))
429 }
430
431 pub fn children<T: ChannelType>(&self, parent: &str) -> Vec<T> {
433 let mut guids: Vec<String> = self
434 .lock_objects()
435 .iter()
436 .filter(|(_, entry)| entry.parent == parent && T::accepts(&entry.interface))
437 .map(|(guid, _)| guid.clone())
438 .collect();
439 guids.sort();
440 guids
441 .into_iter()
442 .map(|guid| T::from_channel(self.channel(&guid)))
443 .collect()
444 }
445
446 pub fn raw_initializer(&self, guid: &str) -> Result<Value, ProtocolError> {
448 self.lock_objects()
449 .get(guid)
450 .map(|entry| entry.initializer.clone())
451 .ok_or_else(|| ProtocolError::UnknownObject(guid.to_string()))
452 }
453
454 pub fn object_count(&self) -> usize {
456 self.lock_objects().len()
457 }
458
459 pub fn subscribe(&self) -> broadcast::Receiver<RawEvent> {
461 self.inner.events.subscribe()
462 }
463
464 pub fn close_reason(&self) -> Option<String> {
466 self.inner
467 .closed
468 .lock()
469 .ok()
470 .and_then(|reason| reason.clone())
471 }
472
473 pub fn set_default_timeout(&self, timeout_ms: Option<f64>) {
476 if let Ok(mut slot) = self.inner.default_timeout.lock() {
477 *slot = timeout_ms;
478 }
479 }
480
481 pub fn default_timeout(&self) -> Option<f64> {
483 self.inner
484 .default_timeout
485 .lock()
486 .ok()
487 .and_then(|slot| *slot)
488 }
489
490 pub async fn call(
492 &self,
493 guid: &str,
494 method: &str,
495 params: Value,
496 ) -> Result<Value, ProtocolError> {
497 self.call_with_timeout(guid, method, params, self.default_timeout())
498 .await
499 }
500
501 pub async fn call_with_timeout(
504 &self,
505 guid: &str,
506 method: &str,
507 params: Value,
508 timeout_ms: Option<f64>,
509 ) -> Result<Value, ProtocolError> {
510 if let Some(reason) = self.close_reason() {
511 return Err(ProtocolError::Closed(reason));
512 }
513 let id = self.inner.next_id.fetch_add(1, Ordering::Relaxed);
514 let (sender, receiver) = oneshot::channel();
515 self.lock_pending().insert(id, sender);
516 let message = json!({
517 "id": id,
518 "guid": guid,
519 "method": method,
520 "params": params,
521 "metadata": match timeout_ms {
522 Some(timeout) => json!({ "timeout": timeout }),
523 None => json!({}),
524 },
525 });
526 let written = {
527 let mut writer = self.inner.writer.lock().await;
528 write_message(&mut *writer, &message).await
529 };
530 if let Err(err) = written {
531 self.lock_pending().remove(&id);
532 return Err(err.into());
533 }
534 match receiver.await {
535 Ok(result) => result,
536 Err(_) => Err(ProtocolError::Closed(
537 self.close_reason()
538 .unwrap_or_else(|| "the response was dropped".to_string()),
539 )),
540 }
541 }
542
543 pub async fn close_input(&self) {
546 use tokio::io::AsyncWriteExt;
547 let mut writer = self.inner.writer.lock().await;
548 let _ = writer.shutdown().await;
549 *writer = Box::new(tokio::io::sink());
552 drop(writer);
553 if let Ok(mut closed) = self.inner.closed.lock() {
554 closed.get_or_insert_with(|| "the connection was closed".to_string());
555 }
556 }
557
558 fn lock_objects(&self) -> std::sync::MutexGuard<'_, HashMap<String, ObjectEntry>> {
559 self.inner
560 .objects
561 .lock()
562 .unwrap_or_else(std::sync::PoisonError::into_inner)
563 }
564
565 fn lock_pending(&self) -> std::sync::MutexGuard<'_, Pending> {
566 self.inner
567 .pending
568 .lock()
569 .unwrap_or_else(std::sync::PoisonError::into_inner)
570 }
571
572 fn mark_closed(&self, reason: String) {
573 if let Ok(mut closed) = self.inner.closed.lock() {
574 closed.get_or_insert(reason.clone());
575 }
576 let pending: Vec<_> = self.lock_pending().drain().collect();
577 for (_, sender) in pending {
578 let _ = sender.send(Err(ProtocolError::Closed(reason.clone())));
579 }
580 }
581
582 pub(crate) fn dispatch(&self, message: Value) {
584 if let Some(id) = message.get("id").and_then(Value::as_u64) {
585 let Some(sender) = self.lock_pending().remove(&id) else {
586 return;
587 };
588 let outcome = match message.get("error") {
589 Some(error) => Err(remote_error(error)),
590 None => Ok(message.get("result").cloned().unwrap_or(Value::Null)),
591 };
592 let _ = sender.send(outcome);
593 return;
594 }
595 let guid = message
596 .get("guid")
597 .and_then(Value::as_str)
598 .unwrap_or_default()
599 .to_string();
600 let method = message
601 .get("method")
602 .and_then(Value::as_str)
603 .unwrap_or_default()
604 .to_string();
605 let params = message.get("params").cloned().unwrap_or(Value::Null);
606 match method.as_str() {
607 "__create__" => {
608 let child = params["guid"].as_str().unwrap_or_default().to_string();
609 let entry = ObjectEntry {
610 interface: params["type"].as_str().unwrap_or_default().to_string(),
611 initializer: params.get("initializer").cloned().unwrap_or(json!({})),
612 parent: guid,
613 };
614 self.lock_objects().insert(child, entry);
615 }
616 "__adopt__" => {
617 if let Some(child) = params["guid"].as_str() {
618 if let Some(entry) = self.lock_objects().get_mut(child) {
619 entry.parent = guid;
620 }
621 }
622 }
623 "__dispose__" => self.dispose(&guid),
624 _ => {
625 self.track_live_state(&guid, &method, ¶ms);
626 let _ = self.inner.events.send(RawEvent {
627 guid,
628 method,
629 params,
630 });
631 }
632 }
633 }
634
635 fn track_live_state(&self, guid: &str, method: &str, params: &Value) {
639 let mut objects = self.lock_objects();
640 let Some(entry) = objects.get_mut(guid) else {
641 return;
642 };
643 if entry.interface != "Frame" {
644 return;
645 }
646 let Some(initializer) = entry.initializer.as_object_mut() else {
647 return;
648 };
649 match method {
650 "navigated" if params.get("error").is_none() => {
651 for key in ["url", "name"] {
652 if let Some(value) = params.get(key) {
653 initializer.insert(key.to_string(), value.clone());
654 }
655 }
656 }
657 "loadstate" => {
658 let states = initializer.entry("loadStates").or_insert_with(|| json!([]));
659 let Some(states) = states.as_array_mut() else {
660 return;
661 };
662 if let Some(added) = params.get("add") {
663 if !states.contains(added) {
664 states.push(added.clone());
665 }
666 }
667 if let Some(removed) = params.get("remove") {
668 states.retain(|state| state != removed);
669 }
670 }
671 _ => {}
672 }
673 }
674
675 fn dispose(&self, guid: &str) {
676 let mut objects = self.lock_objects();
677 let mut stack = vec![guid.to_string()];
678 while let Some(current) = stack.pop() {
679 objects.remove(¤t);
680 stack.extend(
681 objects
682 .iter()
683 .filter(|(_, entry)| entry.parent == current)
684 .map(|(child, _)| child.clone()),
685 );
686 }
687 }
688}
689
690fn remote_error(error: &Value) -> ProtocolError {
691 let details = error.get("error").unwrap_or(error);
692 let text = |key: &str| details.get(key).and_then(Value::as_str).map(str::to_string);
693 ProtocolError::Remote {
694 name: text("name").unwrap_or_else(|| "Error".to_string()),
695 message: text("message").unwrap_or_else(|| details.to_string()),
696 stack: text("stack"),
697 }
698}
699
700#[cfg(test)]
701mod tests {
702 use super::*;
703 use crate::playwright::protocol::{
704 BrowserContext, ElementHandle, JSHandle, Page, PageEvent, PageInitializer,
705 };
706
707 fn offline() -> Connection {
708 let (client, _server) = tokio::io::duplex(1024);
709 let (reader, writer) = tokio::io::split(client);
710 Connection::new(reader, writer)
711 }
712
713 fn create(connection: &Connection, parent: &str, interface: &str, guid: &str, init: Value) {
714 connection.dispatch(json!({
715 "guid": parent,
716 "method": "__create__",
717 "params": { "type": interface, "guid": guid, "initializer": init },
718 }));
719 }
720
721 #[tokio::test]
722 async fn references_resolve_to_typed_objects_and_check_the_interface() {
723 let connection = offline();
724 create(
725 &connection,
726 "",
727 "ElementHandle",
728 "handle@1",
729 json!({ "preview": "<a>" }),
730 );
731
732 let as_element: ElementHandle = connection.object(&Ref::new("handle@1")).unwrap();
733 assert_eq!(as_element.guid(), "handle@1");
734 let as_js: JSHandle = connection.object(&Ref::new("handle@1")).unwrap();
736 assert_eq!(as_js.initializer().unwrap().preview, "<a>");
737 let wrong = connection.object::<Page>(&Ref::new("handle@1"));
738 assert!(matches!(wrong, Err(ProtocolError::WrongInterface { .. })));
739 let missing = connection.object::<Page>(&Ref::new("nope"));
740 assert!(matches!(missing, Err(ProtocolError::UnknownObject(_))));
741 }
742
743 #[tokio::test]
744 async fn dispose_drops_children_and_adopt_moves_them() {
745 let connection = offline();
746 create(&connection, "", "BrowserContext", "context@1", json!({}));
747 create(&connection, "", "BrowserContext", "context@2", json!({}));
748 create(&connection, "context@1", "Page", "page@1", json!({}));
749 assert_eq!(connection.children::<Page>("context@1").len(), 1);
750
751 connection.dispatch(json!({
752 "guid": "context@2", "method": "__adopt__", "params": { "guid": "page@1" },
753 }));
754 assert!(connection.children::<Page>("context@1").is_empty());
755 assert_eq!(connection.children::<Page>("context@2").len(), 1);
756 assert_eq!(connection.children::<BrowserContext>("").len(), 2);
757
758 connection.dispatch(json!({ "guid": "context@2", "method": "__dispose__", "params": {} }));
759 assert_eq!(connection.object_count(), 1);
760 assert!(connection.raw_initializer("page@1").is_err());
761 }
762
763 #[tokio::test]
764 async fn frame_urls_and_load_states_follow_their_events() {
765 use crate::playwright::protocol::{Frame, LifecycleEvent};
766 let connection = offline();
767 create(
768 &connection,
769 "",
770 "Frame",
771 "frame@1",
772 json!({ "url": "about:blank", "name": "", "loadStates": ["load"] }),
773 );
774 let frame: Frame = connection.object(&Ref::new("frame@1")).unwrap();
775 let event = |method: &str, params: Value| {
776 connection.dispatch(json!({ "guid": "frame@1", "method": method, "params": params }));
777 };
778
779 event("navigated", json!({ "url": "https://a.test/", "name": "" }));
780 event("loadstate", json!({ "remove": "load" }));
781 event("loadstate", json!({ "add": "domcontentloaded" }));
782 event(
784 "navigated",
785 json!({ "url": "https://b.test/", "name": "", "error": "net::ERR" }),
786 );
787
788 let state = frame.initializer().unwrap();
789 assert_eq!(state.url, "https://a.test/");
790 assert_eq!(state.load_states, vec![LifecycleEvent::Domcontentloaded]);
791 }
792
793 #[tokio::test]
794 async fn events_are_typed_and_filtered_by_object() {
795 let connection = offline();
796 create(&connection, "", "Page", "page@1", json!({}));
797 create(&connection, "", "Page", "page@2", json!({}));
798 let page: Page = connection.object(&Ref::new("page@1")).unwrap();
799 let mut events = page.events();
800
801 connection.dispatch(json!({ "guid": "page@2", "method": "close", "params": {} }));
802 connection.dispatch(json!({ "guid": "page@1", "method": "crash", "params": {} }));
803 connection
804 .dispatch(json!({ "guid": "page@1", "method": "brandNew", "params": { "x": 1 } }));
805
806 assert_eq!(events.recv().await.unwrap(), PageEvent::Crash);
807 assert!(matches!(
808 events.recv().await.unwrap(),
809 PageEvent::Unknown { method, .. } if method == "brandNew"
810 ));
811 let _: Result<PageInitializer, _> = page.initializer();
813 }
814
815 #[tokio::test]
816 async fn remote_errors_carry_the_driver_error_name() {
817 let error = remote_error(&json!({
818 "error": { "name": "TimeoutError", "message": "Timeout 5ms exceeded", "stack": "s" }
819 }));
820 assert!(error.is_timeout());
821 assert_eq!(error.to_string(), "TimeoutError: Timeout 5ms exceeded");
822 }
823}