Skip to main content

proxyfor/
state.rs

1use crate::{
2    server::PrintMode,
3    traffic::{wrap_entries, Body, Traffic, TrafficHead},
4};
5
6use anyhow::{anyhow, bail, Result};
7use indexmap::IndexMap;
8use serde::Serialize;
9use serde_json::Value;
10use time::OffsetDateTime;
11use tokio::{sync::broadcast, sync::Mutex};
12use tokio_tungstenite::tungstenite;
13
14#[derive(Debug)]
15pub struct State {
16    print_mode: PrintMode,
17    traffics: Mutex<IndexMap<usize, Traffic>>,
18    traffics_notifier: broadcast::Sender<TrafficHead>,
19    websockets: Mutex<IndexMap<usize, Vec<WebsocketMessage>>>,
20    websockets_notifier: broadcast::Sender<(usize, WebsocketMessage)>,
21}
22
23impl State {
24    pub fn new(print_mode: PrintMode) -> Self {
25        let (traffics_notifier, _) = broadcast::channel(128);
26        let (websockets_notifier, _) = broadcast::channel(64);
27        Self {
28            print_mode,
29            traffics: Default::default(),
30            traffics_notifier,
31            websockets: Default::default(),
32            websockets_notifier,
33        }
34    }
35
36    pub async fn add_traffic(&self, traffic: Traffic) {
37        if !traffic.valid {
38            return;
39        }
40        let mut traffics = self.traffics.lock().await;
41        let id = traffics.len() + 1;
42        let head = traffic.head(id);
43        traffics.insert(id, traffic);
44        let _ = self.traffics_notifier.send(head);
45    }
46
47    pub async fn done_traffic(&self, gid: usize, raw_size: u64) {
48        let mut traffics = self.traffics.lock().await;
49        let Some((id, traffic)) = traffics.iter_mut().find(|(_, v)| v.gid == gid) else {
50            return;
51        };
52
53        traffic.uncompress_res_file().await;
54        traffic.done_res_body(raw_size);
55
56        let head = traffic.head(*id);
57        let _ = self.traffics_notifier.send(head);
58        match self.print_mode {
59            PrintMode::Nothing => {}
60            PrintMode::Oneline => {
61                println!("# {}", traffic.oneline());
62            }
63            PrintMode::Markdown => {
64                println!("{}", traffic.markdown().await);
65            }
66        }
67    }
68
69    pub async fn get_traffic(&self, id: usize) -> Option<Traffic> {
70        let traffics = self.traffics.lock().await;
71        traffics.get(&id).cloned()
72    }
73
74    pub fn subscribe_traffics(&self) -> broadcast::Receiver<TrafficHead> {
75        self.traffics_notifier.subscribe()
76    }
77
78    pub async fn list_heads(&self) -> Vec<TrafficHead> {
79        let traffics = self.traffics.lock().await;
80        traffics
81            .iter()
82            .map(|(id, traffic)| traffic.head(*id))
83            .collect()
84    }
85
86    pub async fn export_traffic(&self, id: usize, format: &str) -> Result<(String, &'static str)> {
87        let traffic = self
88            .get_traffic(id)
89            .await
90            .ok_or_else(|| anyhow!("Not found traffic {id}"))?;
91        traffic.export(format).await
92    }
93
94    pub async fn export_all_traffics(&self, format: &str) -> Result<(String, &'static str)> {
95        let traffics = self.traffics.lock().await;
96        match format {
97            "markdown" => {
98                let output =
99                    futures_util::future::join_all(traffics.iter().map(|(_, v)| v.markdown()))
100                        .await
101                        .into_iter()
102                        .collect::<Vec<String>>()
103                        .join("\n\n");
104                Ok((output, "text/markdown; charset=UTF-8"))
105            }
106            "har" => {
107                let values: Vec<Value> =
108                    futures_util::future::join_all(traffics.iter().map(|(_, v)| v.har_entry()))
109                        .await
110                        .into_iter()
111                        .flatten()
112                        .collect();
113                let json_output = wrap_entries(values);
114                let output = serde_json::to_string_pretty(&json_output)?;
115                Ok((output, "application/json; charset=UTF-8"))
116            }
117            "curl" => {
118                let output = futures_util::future::join_all(traffics.iter().map(|(_, v)| v.curl()))
119                    .await
120                    .into_iter()
121                    .collect::<Vec<String>>()
122                    .join("\n\n");
123                Ok((output, "text/plain; charset=UTF-8"))
124            }
125            "json" => {
126                let values = futures_util::future::join_all(traffics.iter().map(|(_, v)| v.json()))
127                    .await
128                    .into_iter()
129                    .collect::<Vec<Value>>();
130                let output = serde_json::to_string_pretty(&values)?;
131                Ok((output, "application/json; charset=UTF-8"))
132            }
133            "" => {
134                let values = traffics
135                    .iter()
136                    .map(|(id, traffic)| traffic.head(*id))
137                    .collect::<Vec<TrafficHead>>();
138                let output = serde_json::to_string_pretty(&values)?;
139                Ok((output, "application/json; charset=UTF-8"))
140            }
141            _ => bail!("Unsupported format: {}", format),
142        }
143    }
144
145    pub async fn new_websocket(&self) -> usize {
146        let mut websockets = self.websockets.lock().await;
147        let id = websockets.len() + 1;
148        websockets.insert(id, vec![]);
149        id
150    }
151
152    pub async fn add_websocket_error(&self, id: usize, error: String) {
153        let mut websockets = self.websockets.lock().await;
154        let Some(messages) = websockets.get_mut(&id) else {
155            return;
156        };
157        let message = WebsocketMessage::Error(error);
158        messages.push(message.clone());
159        let _ = self.websockets_notifier.send((id, message));
160    }
161
162    pub async fn add_websocket_message(
163        &self,
164        id: usize,
165        message: &tungstenite::Message,
166        server_to_client: bool,
167    ) {
168        let mut websockets = self.websockets.lock().await;
169        let Some(messages) = websockets.get_mut(&id) else {
170            return;
171        };
172        let body = match message {
173            tungstenite::Message::Text(text) => Body::text(text),
174            tungstenite::Message::Binary(bin) => Body::bytes(bin),
175            _ => return,
176        };
177        let message = WebsocketMessage::Data(WebsocketData {
178            create: OffsetDateTime::now_utc(),
179            server_to_client,
180            body,
181        });
182        messages.push(message.clone());
183        let _ = self.websockets_notifier.send((id, message));
184    }
185
186    pub async fn subscribe_websocket(&self, id: usize) -> Option<SubscribedWebSocket> {
187        let websockets = self.websockets.lock().await;
188        let messages = websockets.get(&id)?;
189        Some((messages.to_vec(), self.websockets_notifier.subscribe()))
190    }
191}
192
193pub type SubscribedWebSocket = (
194    Vec<WebsocketMessage>,
195    broadcast::Receiver<(usize, WebsocketMessage)>,
196);
197
198#[derive(Debug, Clone, Serialize)]
199pub enum WebsocketMessage {
200    #[serde(rename = "error")]
201    Error(String),
202    #[serde(rename = "data")]
203    Data(WebsocketData),
204}
205
206#[derive(Debug, Clone, Serialize)]
207pub struct WebsocketData {
208    #[serde(serialize_with = "crate::utils::serialize_datetime")]
209    pub create: OffsetDateTime,
210    pub server_to_client: bool,
211    pub body: Body,
212}