Skip to main content

huawei_dongle_api/
session.rs

1//! Session management and CSRF token handling
2
3use crate::error::{Error, Result};
4use reqwest::Client as HttpClient;
5use std::sync::Arc;
6use tokio::sync::RwLock;
7use tracing::{debug, trace};
8use url::Url;
9
10/// Session state for managing authentication and CSRF tokens
11#[derive(Debug, Clone, Default)]
12pub struct SessionState {
13    /// Current CSRF token
14    pub csrf_token: Option<String>,
15    /// Session cookies are managed by reqwest's cookie store
16    pub is_authenticated: bool,
17    /// Username of the authenticated user
18    pub username: Option<String>,
19    /// Last authentication time
20    pub last_auth_time: Option<chrono::DateTime<chrono::Utc>>,
21}
22
23/// Session manager handles CSRF tokens and authentication state
24#[derive(Debug)]
25pub struct SessionManager {
26    http_client: HttpClient,
27    base_url: Url,
28    state: Arc<RwLock<SessionState>>,
29}
30
31impl SessionManager {
32    pub fn new(http_client: HttpClient, base_url: Url) -> Self {
33        Self {
34            http_client,
35            base_url,
36            state: Arc::new(RwLock::new(SessionState::default())),
37        }
38    }
39
40    /// Get the current CSRF token, fetching one if needed
41    pub async fn get_csrf_token(&self) -> Result<String> {
42        {
43            let state = self.state.read().await;
44            if let Some(ref token) = state.csrf_token {
45                trace!("Using cached CSRF token");
46                return Ok(token.clone());
47            }
48        }
49
50        self.refresh_csrf_token().await
51    }
52
53    /// Refresh the CSRF token by fetching from the token endpoint
54    pub async fn refresh_csrf_token(&self) -> Result<String> {
55        debug!("Fetching new CSRF token from /api/webserver/token");
56
57        match self.try_api_token().await {
58            Ok(token) => {
59                debug!("Successfully fetched token from API endpoint");
60                return Ok(token);
61            }
62            Err(e) => {
63                debug!("API token fetch failed: {}, trying homepage fallback", e);
64            }
65        }
66
67        self.try_homepage_token().await
68    }
69
70    /// Try to get CSRF token from the API endpoint
71    async fn try_api_token(&self) -> Result<String> {
72        let url = self.base_url.join("/api/webserver/token")?;
73        let response = self.http_client.get(url).send().await?;
74
75        if !response.status().is_success() {
76            return Err(Error::session(format!(
77                "Failed to fetch token: HTTP {}",
78                response.status()
79            )));
80        }
81
82        let xml = response.text().await?;
83        trace!("Token response XML: {}", xml);
84
85        let token = self.extract_token_from_xml(&xml)?;
86
87        {
88            let mut state = self.state.write().await;
89            state.csrf_token = Some(token.clone());
90        }
91
92        Ok(token)
93    }
94
95    /// Try to get CSRF token from homepage HTML
96    async fn try_homepage_token(&self) -> Result<String> {
97        debug!("Fetching CSRF token from homepage HTML");
98
99        let response = self.http_client.get(self.base_url.clone()).send().await?;
100
101        if !response.status().is_success() {
102            return Err(Error::session(format!(
103                "Failed to fetch homepage: HTTP {}",
104                response.status()
105            )));
106        }
107
108        let html = response.text().await?;
109        trace!("Homepage HTML length: {} chars", html.len());
110
111        let token = self.extract_token_from_html(&html)?;
112
113        {
114            let mut state = self.state.write().await;
115            state.csrf_token = Some(token.clone());
116        }
117
118        debug!("Successfully extracted token from homepage HTML");
119        Ok(token)
120    }
121
122    fn extract_token_from_xml(&self, xml: &str) -> Result<String> {
123        use quick_xml::events::Event;
124        use quick_xml::Reader;
125
126        let mut reader = Reader::from_str(xml);
127        reader.trim_text(true);
128
129        let mut buf = Vec::new();
130        let mut in_token = false;
131
132        loop {
133            match reader.read_event_into(&mut buf)? {
134                Event::Start(ref e) if e.name().as_ref() == b"token" => {
135                    in_token = true;
136                }
137                Event::Text(e) if in_token => {
138                    let token = e.unescape()?.into_owned();
139                    return Ok(token);
140                }
141                Event::End(ref e) if e.name().as_ref() == b"token" => {
142                    in_token = false;
143                }
144                Event::Eof => break,
145                _ => (),
146            }
147            buf.clear();
148        }
149
150        Err(Error::session("Could not find token in XML response"))
151    }
152
153    /// Extract CSRF token from HTML homepage
154    fn extract_token_from_html(&self, html: &str) -> Result<String> {
155        use scraper::{Html, Selector};
156
157        let document = Html::parse_document(html);
158
159        let meta_selector = Selector::parse(r#"meta[name="csrf_token"]"#)
160            .map_err(|_| Error::session("Invalid CSS selector"))?;
161
162        if let Some(meta_element) = document.select(&meta_selector).next() {
163            if let Some(content) = meta_element.value().attr("content") {
164                if !content.is_empty() {
165                    return Ok(content.to_string());
166                }
167            }
168        }
169
170        let token_selector = Selector::parse(r#"meta[content*="csrf"]"#)
171            .map_err(|_| Error::session("Invalid CSS selector"))?;
172
173        if let Some(meta_element) = document.select(&token_selector).next() {
174            if let Some(content) = meta_element.value().attr("content") {
175                if !content.is_empty() {
176                    return Ok(content.to_string());
177                }
178            }
179        }
180
181        let all_meta_selector =
182            Selector::parse("meta[content]").map_err(|_| Error::session("Invalid CSS selector"))?;
183
184        for meta_element in document.select(&all_meta_selector) {
185            if let Some(content) = meta_element.value().attr("content") {
186                if content.len() > 20 && content.chars().all(|c| c.is_alphanumeric()) {
187                    debug!("Found potential token in meta tag: {}...", &content[..10]);
188                    return Ok(content.to_string());
189                }
190            }
191        }
192
193        Err(Error::session("Could not find CSRF token in HTML"))
194    }
195
196    pub async fn clear_session(&self) {
197        let mut state = self.state.write().await;
198        state.csrf_token = None;
199        state.is_authenticated = false;
200        state.username = None;
201        state.last_auth_time = None;
202        debug!("Session state cleared");
203    }
204
205    pub async fn is_authenticated(&self) -> bool {
206        let state = self.state.read().await;
207        state.is_authenticated
208    }
209
210    /// Mark session as invalidated (e.g., after getting 401)
211    pub async fn invalidate_session(&self) {
212        debug!("Session invalidated, will need to re-authenticate");
213        self.clear_session().await;
214    }
215
216    /// Mark user as authenticated
217    pub async fn mark_authenticated(&self, username: &str) {
218        let mut state = self.state.write().await;
219        state.is_authenticated = true;
220        state.username = Some(username.to_string());
221        state.last_auth_time = Some(chrono::Utc::now());
222        debug!("User '{}' marked as authenticated", username);
223    }
224
225    pub async fn current_username(&self) -> Option<String> {
226        let state = self.state.read().await;
227        state.username.clone()
228    }
229
230    pub async fn last_auth_time(&self) -> Option<chrono::DateTime<chrono::Utc>> {
231        let state = self.state.read().await;
232        state.last_auth_time
233    }
234
235    /// Check if the session is expired (based on time)
236    pub async fn is_session_expired(&self, max_age_minutes: u64) -> bool {
237        let state = self.state.read().await;
238        if let Some(last_auth) = state.last_auth_time {
239            let now = chrono::Utc::now();
240            let age = now.signed_duration_since(last_auth);
241            age.num_minutes() > max_age_minutes as i64
242        } else {
243            true
244        }
245    }
246
247    /// Update CSRF token from response headers if available
248    pub async fn update_token_from_headers(&self, headers: &reqwest::header::HeaderMap) {
249        let token_headers = [
250            "__RequestVerificationToken",
251            "__RequestVerificationTokenone",
252            "__RequestVerificationTokentwo",
253        ];
254
255        for header_name in &token_headers {
256            if let Some(token_value) = headers.get(*header_name) {
257                if let Ok(token_str) = token_value.to_str() {
258                    if !token_str.is_empty() {
259                        let mut state = self.state.write().await;
260                        state.csrf_token = Some(token_str.to_string());
261                        debug!(
262                            "Updated CSRF token from response header {}: {}",
263                            header_name, token_str
264                        );
265                        return;
266                    }
267                }
268            }
269        }
270    }
271}
272
273#[cfg(test)]
274mod tests {
275    use super::*;
276    use reqwest::header::HeaderMap;
277
278    #[tokio::test]
279    async fn test_update_token_from_headers() {
280        let http_client = reqwest::Client::new();
281        let base_url = Url::parse("http://192.168.8.1").unwrap();
282        let session = SessionManager::new(http_client, base_url);
283
284        let mut state = session.state.write().await;
285        state.csrf_token = Some("old_token".to_string());
286        drop(state);
287
288        let mut headers = HeaderMap::new();
289        headers.insert("__RequestVerificationToken", "new_token".parse().unwrap());
290
291        session.update_token_from_headers(&headers).await;
292
293        let state = session.state.read().await;
294        assert_eq!(state.csrf_token, Some("new_token".to_string()));
295    }
296
297    #[tokio::test]
298    async fn test_update_token_from_headers_alternate_names() {
299        let http_client = reqwest::Client::new();
300        let base_url = Url::parse("http://192.168.8.1").unwrap();
301        let session = SessionManager::new(http_client, base_url);
302
303        let mut headers = HeaderMap::new();
304        headers.insert(
305            "__RequestVerificationTokenone",
306            "token_one".parse().unwrap(),
307        );
308        session.update_token_from_headers(&headers).await;
309
310        let state = session.state.read().await;
311        assert_eq!(state.csrf_token, Some("token_one".to_string()));
312        drop(state);
313
314        let mut headers = HeaderMap::new();
315        headers.insert(
316            "__RequestVerificationTokentwo",
317            "token_two".parse().unwrap(),
318        );
319        session.update_token_from_headers(&headers).await;
320
321        let state = session.state.read().await;
322        assert_eq!(state.csrf_token, Some("token_two".to_string()));
323    }
324
325    #[tokio::test]
326    async fn test_no_token_update_when_missing() {
327        let http_client = reqwest::Client::new();
328        let base_url = Url::parse("http://192.168.8.1").unwrap();
329        let session = SessionManager::new(http_client, base_url);
330
331        let mut state = session.state.write().await;
332        state.csrf_token = Some("existing_token".to_string());
333        drop(state);
334
335        let headers = HeaderMap::new();
336        session.update_token_from_headers(&headers).await;
337
338        let state = session.state.read().await;
339        assert_eq!(state.csrf_token, Some("existing_token".to_string()));
340    }
341}