1use std::future::Future;
2use std::sync::Arc;
3use base64::Engine;
4use base64::prelude::BASE64_STANDARD;
5use serde::Serialize;
6use tokio_jrpc::ClientHandle;
7use ulid::Ulid;
8use crate::context_menu::{ContextMenuNodeMulti, ContextMenuNodeSingle};
9use crate::error::Error;
10use crate::inspector::MessageTab;
11use crate::proxy_server::{ConnectContext, ConnectResult};
12use crate::http2::Http2Event;
13use crate::overview::OverviewNode;
14use crate::runtime::state::{ConnectHandler, ExtensionState};
15use crate::sessions::{SessionEntry, SessionInfo};
16use crate::tls::TlsEvent;
17use crate::websocket::WebSocketMessage;
18
19#[derive(Clone)]
24pub struct ExtensionHandle {
25 client: ClientHandle,
26 state: Arc<ExtensionState>,
27}
28
29impl ExtensionHandle {
30 pub(crate) fn new(client: ClientHandle, state: Arc<ExtensionState>) -> Self {
31 Self { client, state }
32 }
33
34 pub async fn shutdown(&self) {
36 self.state.shutdown_handle().shutdown().await;
37 }
38
39 pub async fn write_text_to_clipboard(&self, text: &str) -> Result<(), Error> {
43 self.client
44 .request("clipboard/write_text", ClipboardWriteTextParams { text }).await
45 .map_err(Error::from)
46 }
47
48 pub async fn extend_context_menu_single<N: Into<ContextMenuNodeSingle>>(&self, node: N) -> Result<(), Error> {
55 let node = node.into();
56 let handler_pairs = node.extract_handlers();
57 let handler_ids: Vec<String> = handler_pairs.iter().map(|(id, _)| id.clone()).collect();
58
59 self.state.extend_context_menu_single_handlers(handler_pairs).await;
60
61 if let Err(err) = self.client.request::<()>("context_menu/extend_single", &node).await {
62 self.state.remove_context_menu_single_handler(&handler_ids).await;
63 return Err(err.into());
64 }
65 Ok(())
66 }
67
68 pub async fn remove_context_menu_item_single(&self, item_id: &str) -> Result<(), Error> {
72 self.client.request::<()>(
73 "context_menu/remove_item_single",
74 ContextMenuItemRef { item_id },
75 ).await?;
76
77 self.state.remove_context_menu_single_handler(&[item_id]).await;
78 Ok(())
79 }
80
81 pub async fn extend_context_menu_multi<N: Into<ContextMenuNodeMulti>>(&self, node: N) -> Result<(), Error> {
88 let node = node.into();
89 let handler_pairs = node.extract_handlers();
90 let handler_ids: Vec<String> = handler_pairs.iter().map(|(id, _)| id.clone()).collect();
91
92 self.state.extend_context_menu_multi_handlers(handler_pairs).await;
93
94 if let Err(err) = self.client.request::<()>("context_menu/extend_multi", &node).await {
95 self.state.remove_context_menu_multi_handler(&handler_ids).await;
96 return Err(err.into());
97 }
98 Ok(())
99 }
100
101 pub async fn remove_context_menu_item_multi(&self, item_id: &str) -> Result<(), Error> {
105 self.client.request::<()>(
106 "context_menu/remove_item_multi",
107 ContextMenuItemRef { item_id },
108 ).await?;
109
110 self.state.remove_context_menu_multi_handler(&[item_id]).await;
111 Ok(())
112 }
113
114 pub async fn extend_overview<N: Into<OverviewNode>>(&self, node: N) -> Result<(), Error> {
121 let node = node.into();
122 let handler_pairs = node.extract_handlers();
123 let handler_ids: Vec<String> = handler_pairs.iter().map(|(id, _)| id.clone()).collect();
124
125 self.state.extend_overview_handlers(handler_pairs).await;
126
127 if let Err(err) = self.client.request::<()>("overview/extend", &node).await {
128 self.state.remove_overview_handler(&handler_ids).await;
129 return Err(err.into());
130 }
131 Ok(())
132 }
133
134 pub async fn remove_overview_field(&self, field_id: &str) -> Result<(), Error> {
138 self.client.request::<()>(
139 "overview/remove_field",
140 OverviewFieldRef { field_id }
141 ).await?;
142
143 self.state.remove_overview_handler(&[field_id]).await;
144 Ok(())
145 }
146
147 pub async fn add_request_tab(&self, tab: MessageTab) -> Result<(), Error> {
151 let (tab_id, handlers) = tab.extract_handlers();
152
153 self.state.insert_request_tab_handlers(tab_id.clone(), handlers).await;
154
155 if let Err(err) = self.client.request::<()>("inspector/add_request_tab", &tab).await {
156 self.state.remove_request_tab_handlers(&tab_id).await;
157 return Err(err.into());
158 }
159 Ok(())
160 }
161
162 pub async fn remove_request_tab(&self, tab_id: &str) -> Result<(), Error> {
166 self.client.request::<()>(
167 "inspector/remove_request_tab",
168 MessageTabRef { tab_id },
169 ).await?;
170
171 self.state.remove_request_tab_handlers(tab_id).await;
172 Ok(())
173 }
174
175 pub async fn add_response_tab(&self, tab: MessageTab) -> Result<(), Error> {
179 let (tab_id, handlers) = tab.extract_handlers();
180
181 self.state.insert_response_tab_handlers(tab_id.clone(), handlers).await;
182
183 if let Err(err) = self.client.request::<()>("inspector/add_response_tab", &tab).await {
184 self.state.remove_response_tab_handlers(&tab_id).await;
185 return Err(err.into());
186 }
187 Ok(())
188 }
189
190 pub async fn remove_response_tab(&self, tab_id: &str) -> Result<(), Error> {
194 self.client.request::<()>(
195 "inspector/remove_response_tab",
196 MessageTabRef { tab_id },
197 ).await?;
198
199 self.state.remove_response_tab_handlers(tab_id).await;
200 Ok(())
201 }
202
203 pub async fn add_connect_handler<F, Fut>(&self, handler_id: &str, handler: F) -> Result<(), Error>
210 where
211 F: Fn(ConnectContext, ExtensionHandle) -> Fut + Send + Sync + 'static,
212 Fut: Future<Output = Result<ConnectResult, Error>> + Send + 'static,
213 {
214 let id = handler_id.to_owned();
215 let handler: ConnectHandler = Arc::new(move |ctx, handle| {
216 let fut = handler(ctx, handle);
217 Box::pin(async move { fut.await.map_err(Error::into_jrpc) })
218 });
219
220 self.state.insert_connect_handlers(id.clone(), handler).await;
221
222 let add_result = self.client.request::<()>(
223 "proxy_server/add_connect_handler",
224 ConnectHandlerRef { handler_id }
225 ).await;
226
227 if let Err(err) = add_result {
228 self.state.remove_connect_handler(&id).await;
229 return Err(err.into());
230 }
231 Ok(())
232 }
233
234 pub async fn remove_connect_handler(&self, handler_id: &str) -> Result<(), Error> {
238 self.client.request::<()>(
239 "proxy_server/remove_connect_handler",
240 ConnectHandlerRef { handler_id },
241 ).await?;
242
243 self.state.remove_connect_handler(&handler_id).await;
244 Ok(())
245 }
246
247 pub async fn list_sessions(&self) -> Result<Vec<SessionInfo>, Error> {
251 self.client
252 .request("sessions/list", ()).await
253 .map_err(Error::from)
254 }
255
256 pub async fn get_listener_port(&self, session_id: Ulid) -> Result<Option<u16>, Error> {
260 self.client
261 .request("sessions/get_listener_port", SessionRef { session_id }).await
262 .map_err(Error::from)
263 }
264
265 pub async fn get_session_entry_ids(&self, session_id: Ulid) -> Result<Option<Vec<Ulid>>, Error> {
269 self.client
270 .request("sessions/get_entry_ids", SessionRef { session_id }).await
271 .map_err(Error::from)
272 }
273
274 pub async fn get_session_entry(&self, session_id: Ulid, entry_id: Ulid) -> Result<Option<SessionEntry>, Error> {
278 self.client
279 .request("sessions/get_entry", SessionEntryRef { session_id, entry_id }).await
280 .map_err(Error::from)
281 }
282
283 pub async fn get_request_body_as_text(&self, session_id: Ulid, entry_id: Ulid) -> Result<Option<String>, Error> {
287 self.client
288 .request(
289 "sessions/get_request_body",
290 GetBodyParams { session_id, entry_id, encoding: BodyEncoding::Text }
291 )
292 .await
293 .map_err(Error::from)
294 }
295
296 pub async fn get_request_body_as_bytes(&self, session_id: Ulid, entry_id: Ulid) -> Result<Option<Vec<u8>>, Error> {
300 self.client
301 .request::<Option<String>>(
302 "sessions/get_request_body",
303 GetBodyParams { session_id, entry_id, encoding: BodyEncoding::Base64 },
304 )
305 .await?
306 .map(|body_text| BASE64_STANDARD.decode(body_text).map_err(Error::new))
307 .transpose()
308 }
309
310 pub async fn get_response_body_as_text(&self, session_id: Ulid, entry_id: Ulid) -> Result<Option<String>, Error> {
314 self.client
315 .request(
316 "sessions/get_response_body",
317 GetBodyParams { session_id, entry_id, encoding: BodyEncoding::Text },
318 )
319 .await
320 .map_err(Error::from)
321 }
322
323 pub async fn get_response_body_as_bytes(&self, session_id: Ulid, entry_id: Ulid) -> Result<Option<Vec<u8>>, Error> {
327 self.client
328 .request::<Option<String>>(
329 "sessions/get_response_body",
330 GetBodyParams { session_id, entry_id, encoding: BodyEncoding::Base64 },
331 )
332 .await?
333 .map(|body_text| BASE64_STANDARD.decode(body_text).map_err(Error::new))
334 .transpose()
335 }
336
337 pub async fn get_websocket_messages(&self, session_id: Ulid, entry_id: Ulid) -> Result<Option<Vec<WebSocketMessage>>, Error> {
341 self.client
342 .request(
343 "sessions/get_websocket_messages",
344 SessionEntryRef { session_id, entry_id },
345 )
346 .await
347 .map_err(Error::from)
348 }
349
350 pub async fn get_tls_connection(&self, connection_id: Ulid) -> Result<Option<Vec<TlsEvent>>, Error> {
354 self.client
355 .request("tls/get_connection", TlsConnectionRef { connection_id }).await
356 .map_err(Error::from)
357 }
358
359 pub async fn get_http2_stream_ids(&self, connection_id: Ulid) -> Result<Option<Vec<u32>>, Error> {
363 self.client
364 .request("http2/get_stream_ids", Http2ConnectionRef { connection_id }).await
365 .map_err(Error::from)
366 }
367
368 pub async fn get_http2_stream(&self, connection_id: Ulid, stream_id: u32) -> Result<Option<Vec<Http2Event>>, Error> {
372 self.client
373 .request("http2/get_stream", Http2StreamRef { connection_id, stream_id }).await
374 .map_err(Error::from)
375 }
376}
377
378#[derive(Serialize)]
379struct ClipboardWriteTextParams<'a> {
380 text: &'a str,
381}
382
383#[derive(Serialize)]
384#[serde(rename_all = "camelCase")]
385struct ContextMenuItemRef<'a> {
386 item_id: &'a str,
387}
388
389#[derive(Serialize)]
390#[serde(rename_all = "camelCase")]
391struct OverviewFieldRef<'a> {
392 field_id: &'a str,
393}
394
395#[derive(Serialize)]
396#[serde(rename_all = "camelCase")]
397struct MessageTabRef<'a> {
398 tab_id: &'a str,
399}
400
401#[derive(Serialize)]
402#[serde(rename_all = "camelCase")]
403struct ConnectHandlerRef<'a> {
404 handler_id: &'a str,
405}
406
407#[derive(Serialize)]
408#[serde(rename_all = "camelCase")]
409struct SessionRef {
410 session_id: Ulid,
411}
412
413#[derive(Serialize)]
414#[serde(rename_all = "camelCase")]
415struct SessionEntryRef {
416 session_id: Ulid,
417 entry_id: Ulid,
418}
419
420#[derive(Serialize)]
421#[serde(rename_all = "camelCase")]
422struct GetBodyParams {
423 session_id: Ulid,
424 entry_id: Ulid,
425 encoding: BodyEncoding,
426}
427
428#[derive(Serialize)]
429#[serde(rename_all = "snake_case")]
430pub enum BodyEncoding {
431 Text,
432 #[serde(rename = "base64")]
433 Base64,
434}
435
436#[derive(Serialize)]
437#[serde(rename_all = "camelCase")]
438struct TlsConnectionRef {
439 connection_id: Ulid,
440}
441
442#[derive(Serialize)]
443#[serde(rename_all = "camelCase")]
444struct Http2ConnectionRef {
445 connection_id: Ulid,
446}
447
448#[derive(Serialize)]
449#[serde(rename_all = "camelCase")]
450pub struct Http2StreamRef {
451 pub connection_id: Ulid,
452 pub stream_id: u32,
453}