Skip to main content

cdp_server/
server.rs

1// @trace REQ-CDS-001 [entity:CdpServer]
2// @trace REQ-CDS-002 [entity:CdpTarget]
3// @trace REQ-CDS-003 [entity:CdpSessionGeneric]
4// @trace REQ-CDS-007 [entity:CdpServer]
5// CdpServer main event loop: TCP accept, HTTP discovery, WS upgrade,
6// command routing, target management.
7
8use std::collections::HashMap;
9use std::io::Read;
10use std::net::{TcpListener, TcpStream};
11use std::sync::{Arc, Mutex};
12use std::time::Duration;
13
14use tungstenite::accept;
15
16use crate::bao_event::ConsoleMessage;
17use crate::event::EventBroadcaster;
18use crate::registry::SharedRegistry;
19use crate::session::{CdpSession, ReplayStream};
20use crate::transport::{self, TargetInfo};
21use crate::{EventSender, ServerConfig, TargetProvider};
22
23pub struct CdpServer {
24    config: ServerConfig,
25    registry: SharedRegistry,
26    target_provider: Option<Arc<dyn TargetProvider>>,
27    broadcaster: Arc<EventBroadcaster>,
28    sessions: Arc<Mutex<HashMap<String, Arc<crate::session::SessionHandle>>>>,
29    /// Receiver for typed console messages forwarded from servo delegates.
30    /// Each message is either a structured CDP event (ConsoleMessage::Event)
31    /// or a plain log (ConsoleMessage::Log).
32    console_rx: Option<std::sync::mpsc::Receiver<ConsoleMessage>>,
33}
34
35impl CdpServer {
36    /// Create CdpServer with an empty domain registry.
37    /// For production use, prefer `with_registry()` with a pre-built registry
38    /// (e.g. `DomainRegistry<DomainDispatch>` for enum dispatch).
39    pub fn new(config: ServerConfig) -> Self {
40        let registry: Arc<crate::DomainRegistry<crate::EmptyHandler>> =
41            Arc::new(crate::DomainRegistry::new());
42        Self::with_registry(config, registry)
43    }
44
45    /// Create CdpServer with a pre-built registry (e.g. DomainRegistry<DomainDispatch>
46    /// for enum dispatch). Any `Arc<R>` where `R: RegistryDispatch` is accepted.
47    pub fn with_registry<R: crate::RegistryDispatch + 'static>(
48        config: ServerConfig,
49        registry: Arc<R>,
50    ) -> Self {
51        let sessions = Arc::new(Mutex::new(HashMap::new()));
52        let broadcaster = Arc::new(EventBroadcaster::new(Arc::clone(&sessions)));
53        CdpServer {
54            config,
55            registry,
56            target_provider: None,
57            broadcaster,
58            sessions,
59            console_rx: None,
60        }
61    }
62
63    pub fn registry(&self) -> &SharedRegistry {
64        &self.registry
65    }
66
67    pub fn broadcaster(&self) -> Arc<EventBroadcaster> {
68        Arc::clone(&self.broadcaster)
69    }
70
71    pub fn set_target_provider(&mut self, provider: Arc<dyn TargetProvider>) {
72        self.target_provider = Some(provider);
73    }
74
75    /// Set the typed console message receiver. Messages are ConsoleMessage
76    /// variants forwarded from servo's show_console_message callbacks.
77    pub fn set_console_receiver(&mut self, rx: std::sync::mpsc::Receiver<ConsoleMessage>) {
78        self.console_rx = Some(rx);
79    }
80
81    pub fn port(&self) -> u16 {
82        self.config.port
83    }
84
85    pub fn ws_url_for_target(&self, target_id: &str) -> String {
86        format!(
87            "ws://{}:{}/devtools/page/{}",
88            self.config.host, self.config.port, target_id
89        )
90    }
91
92    /// Main event loop. Blocks until shutdown.
93    pub fn run(&mut self) -> Result<(), String> {
94        let addr = format!("{}:{}", self.config.host, self.config.port);
95        let listener = TcpListener::bind(&addr).map_err(|e| format!("bind: {}", e))?;
96        listener
97            .set_nonblocking(true)
98            .map_err(|e| format!("nonblocking: {}", e))?;
99
100        log::info!(
101            "CDP listening on ws://{}:{}",
102            self.config.host,
103            self.config.port
104        );
105
106        loop {
107            // Drain session events (not used currently, but placeholder for future command channel).
108            self.check_session_timeouts();
109
110            // Accept new connections.
111            match listener.accept() {
112                Ok((stream, _addr)) => {
113                    self.handle_connection(stream);
114                }
115                Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {}
116                Err(e) => log::warn!("CDP accept error: {}", e),
117            }
118
119            // Process existing sessions. The session map lock is released
120            // BEFORE processing: command dispatch may synchronously emit
121            // events through the EventBroadcaster, which locks the same map
122            // (deadlock if held here). Events land in per-session outboxes
123            // and are drained into the socket here, under the session lock.
124            let mut to_remove = Vec::new();
125            {
126                let session_list: Vec<_> = {
127                    match self.sessions.lock() {
128                        Ok(sessions) => sessions
129                            .iter()
130                            .map(|(id, h)| (id.clone(), Arc::clone(h)))
131                            .collect(),
132                        Err(_) => Vec::new(),
133                    }
134                };
135                for (id, handle) in session_list {
136                    let mut session = match handle.session.lock() {
137                        Ok(s) => s,
138                        Err(_) => continue,
139                    };
140                    let event_sender: Box<dyn EventSender> = self.broadcaster.sender();
141                    if session
142                        .process(&self.registry, event_sender.as_ref())
143                        .is_err()
144                    {
145                        let domains = session.enabled_domains();
146                        let sid = session.session_id().to_string();
147                        session.begin_close();
148                        drop(session);
149                        to_remove.push(id);
150                        self.registry.notify_session_destroyed(&domains, &sid);
151                        continue;
152                    }
153                    // Drain queued events into the socket (gating applied
154                    // here, where the session state is readable).
155                    let drained: Vec<_> = match handle.outbox.lock() {
156                        Ok(mut outbox) => outbox.drain(..).collect(),
157                        Err(_) => Vec::new(),
158                    };
159                    for entry in drained {
160                        let deliver = if entry.browser_only {
161                            session.is_browser_session()
162                        } else {
163                            session.is_browser_session()
164                                || session.has_domain_enabled(&entry.domain)
165                        };
166                        if deliver {
167                            let _ = session.send_text(&entry.json);
168                        }
169                    }
170                }
171            }
172
173            for id in to_remove {
174                if let Ok(mut sessions) = self.sessions.lock() {
175                    if let Some(handle) = sessions.remove(&id) {
176                        if let Ok(mut s) = handle.session.lock() {
177                            s.finalize();
178                        }
179                    }
180                }
181            }
182
183            // Drain typed console messages from servo delegates and broadcast as CDP events.
184            // ConsoleMessage::Event variants are routed to domain-specific events via BaoEvent::broadcast().
185            // ConsoleMessage::Log variants are forwarded as Runtime.consoleAPICalled + Log.entryAdded.
186            if let Some(ref rx) = self.console_rx {
187                while let Ok(msg) = rx.try_recv() {
188                    match msg {
189                        ConsoleMessage::Event(event) => {
190                            event.broadcast(&*self.broadcaster);
191                        }
192                        ConsoleMessage::Log { level, text } => {
193                            self.broadcaster.send_event(
194                                "Runtime.consoleAPICalled",
195                                serde_json::json!({
196                                    "type": match level.as_str() {
197                                        "debug" => "debug",
198                                        "info" => "info",
199                                        "warning" => "warning",
200                                        "error" => "error",
201                                        "verbose" => "verbose",
202                                        _ => "log",
203                                    },
204                                    "args": [serde_json::json!(text)],
205                                    "timestamp": std::time::SystemTime::now()
206                                        .duration_since(std::time::UNIX_EPOCH)
207                                        .unwrap_or_default()
208                                        .as_millis() as f64,
209                                }),
210                            );
211                            self.broadcaster.send_event(
212                                "Log.entryAdded",
213                                serde_json::json!({
214                                    "entry": {
215                                        "source": "javascript",
216                                        "level": level,
217                                        "text": text,
218                                        "timestamp": std::time::SystemTime::now()
219                                            .duration_since(std::time::UNIX_EPOCH)
220                                            .unwrap_or_default()
221                                            .as_millis() as f64,
222                                    }
223                                }),
224                            );
225                        }
226                    }
227                }
228            }
229
230            std::thread::sleep(Duration::from_millis(10));
231        }
232    }
233
234    fn handle_connection(&self, mut stream: TcpStream) {
235        let mut buf = [0u8; 8192];
236        stream.set_nonblocking(false).ok();
237        let n = match stream.read(&mut buf) {
238            Ok(n) if n > 0 => n,
239            _ => return,
240        };
241        let request = match std::str::from_utf8(&buf[..n]) {
242            Ok(s) => s,
243            Err(_) => return,
244        };
245
246        // Check for close/activate/new before general handling.
247        if let Some(target_id) = transport::parse_close_request(request) {
248            if let Some(ref provider) = self.target_provider {
249                match provider.close_target(&target_id) {
250                    Ok(()) => {
251                        transport::respond_json(
252                            &mut stream,
253                            &serde_json::json!({"success": true, "targetId": target_id}),
254                        );
255                        // Broadcast Target.targetDestroyed event.
256                        self.broadcaster.send_event(
257                            "Target.targetDestroyed",
258                            serde_json::json!({"targetId": target_id}),
259                        );
260                    }
261                    Err(e) => {
262                        transport::respond_raw(&mut stream, &format!("500 {}", e));
263                    }
264                }
265            } else {
266                transport::respond_raw(&mut stream, "500 No target provider");
267            }
268            return;
269        }
270
271        if let Some(target_id) = transport::parse_activate_request(request) {
272            if let Some(ref provider) = self.target_provider {
273                match provider.activate_target(&target_id) {
274                    Ok(()) => transport::respond_raw(&mut stream, "Target activated"),
275                    Err(e) => transport::respond_raw(&mut stream, &format!("500 {}", e)),
276                }
277            }
278            return;
279        }
280
281        if let Some(url) = transport::parse_new_request(request) {
282            if let Some(ref provider) = self.target_provider {
283                match provider.create_target(&url) {
284                    Ok(info) => {
285                        let json = serde_json::to_value(&info).unwrap_or_default();
286                        transport::respond_json(&mut stream, &json);
287                    }
288                    Err(e) => {
289                        transport::respond_raw(&mut stream, &format!("500 {}", e));
290                    }
291                }
292            }
293            return;
294        }
295
296        // GET /json/version and /json/list
297        if request.starts_with("GET /json/version")
298            || (request.starts_with("GET /json") && !request.starts_with("GET /json/"))
299        {
300            let targets = self.get_target_list();
301            transport::handle_http_request(&mut stream, request, &self.config, &targets);
302            return;
303        }
304
305        // WebSocket upgrade.
306        if request.contains("Upgrade: websocket") || request.contains("upgrade: websocket") {
307            let (target_id, is_browser) =
308                if let Some(rest) = request.strip_prefix("GET /devtools/page/") {
309                    (rest.split(' ').next().unwrap_or("").to_string(), false)
310                } else if request.starts_with("GET /devtools/browser") {
311                    ("__browser__".to_string(), true)
312                } else {
313                    return;
314                };
315
316            let replay = ReplayStream::new(stream, buf[..n].to_vec());
317            let ws = match accept(replay) {
318                Ok(ws) => ws,
319                Err(e) => {
320                    log::warn!("CDP WebSocket accept error: {}", e);
321                    return;
322                }
323            };
324
325            let session_id = generate_session_id();
326            let session = CdpSession::new(session_id.clone(), target_id, ws, is_browser);
327            let session_count = self.sessions.lock().map(|m| m.len()).unwrap_or(0);
328            if session_count >= self.config.max_sessions {
329                log::warn!("CDP max sessions reached, rejecting");
330                return;
331            }
332            if let Ok(mut sessions) = self.sessions.lock() {
333                sessions.insert(session_id, crate::session::SessionHandle::new(session));
334            }
335        } else {
336            transport::respond_raw(
337                &mut stream,
338                "HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\n\r\n",
339            );
340        }
341    }
342
343    fn get_target_list(&self) -> Vec<TargetInfo> {
344        if let Some(ref provider) = self.target_provider {
345            provider.list_targets()
346        } else {
347            Vec::new()
348        }
349    }
350
351    fn check_session_timeouts(&self) {
352        // Placeholder for future session timeout management.
353    }
354}
355
356fn generate_session_id() -> String {
357    use std::time::{SystemTime, UNIX_EPOCH};
358    let d = SystemTime::now()
359        .duration_since(UNIX_EPOCH)
360        .unwrap_or_default();
361    let ns = d.as_nanos() as u64;
362    format!("{:016x}", ns ^ (ns >> 17) ^ (ns >> 35))
363}
364
365// ---------------------------------------------------------------------------
366// ยง Tests
367// ---------------------------------------------------------------------------
368
369#[cfg(test)]
370mod tests {
371    use super::*;
372    use crate::bao_event::BaoEvent;
373
374    #[test]
375    fn cdp_server_config_stores_host_port_browser_name() {
376        let config = ServerConfig {
377            host: "127.0.0.1".into(),
378            port: 9222,
379            browser_name: "Bao/0.1.0".into(),
380            ..Default::default()
381        };
382        let server = CdpServer::new(config);
383        assert_eq!(server.port(), 9222);
384    }
385
386    #[test]
387    fn server_config_default_values() {
388        let config = ServerConfig::default();
389        assert_eq!(config.host, "127.0.0.1");
390        assert_eq!(config.port, 9222);
391        assert_eq!(config.http_timeout_seconds, 30);
392        assert_eq!(config.max_sessions, 100);
393        assert_eq!(config.browser_name, "Bao/0.1.0");
394        assert_eq!(config.protocol_version, "1.3");
395        assert!(config.user_agent.is_none());
396        assert!(config.v8_version.is_none());
397        assert!(config.webkit_version.is_none());
398    }
399
400    #[test]
401    fn server_config_builder_pattern() {
402        let config = ServerConfig::builder()
403            .host("0.0.0.0")
404            .port(9333)
405            .http_timeout_seconds(60)
406            .max_sessions(200)
407            .browser_name("TestBrowser/1.0")
408            .user_agent("TestAgent")
409            .v8_version("12.0")
410            .webkit_version("602.1")
411            .build();
412        assert_eq!(config.host, "0.0.0.0");
413        assert_eq!(config.port, 9333);
414        assert_eq!(config.http_timeout_seconds, 60);
415        assert_eq!(config.max_sessions, 200);
416        assert_eq!(config.browser_name, "TestBrowser/1.0");
417        assert_eq!(config.user_agent, Some("TestAgent".into()));
418        assert_eq!(config.v8_version, Some("12.0".into()));
419        assert_eq!(config.webkit_version, Some("602.1".into()));
420    }
421
422    #[test]
423    fn ws_url_format_contains_host_port() {
424        let config = ServerConfig {
425            host: "127.0.0.1".into(),
426            port: 9222,
427            ..Default::default()
428        };
429        let server = CdpServer::new(config);
430        let ws_url = server.ws_url_for_target("abc123");
431        assert!(ws_url.starts_with("ws://127.0.0.1:9222/devtools/page/"));
432        assert!(ws_url.ends_with("abc123"));
433    }
434
435    #[test]
436    fn generate_session_id_format() {
437        let id = generate_session_id();
438        assert_eq!(id.len(), 16);
439        assert!(id.chars().all(|c| c.is_ascii_hexdigit()));
440    }
441
442    #[test]
443    fn cdp_server_has_registry_and_broadcaster() {
444        let server = CdpServer::new(ServerConfig::default());
445        let _registry = server.registry();
446        let _broadcaster = server.broadcaster();
447    }
448
449    // --- Console receiver tests (REQ-CDP-007) ---
450
451    #[test]
452    fn cdp_server_default_has_no_console_receiver() {
453        let server = CdpServer::new(ServerConfig::default());
454        assert!(server.console_rx.is_none());
455    }
456
457    #[test]
458    fn cdp_server_set_console_receiver_stores_receiver() {
459        let mut server = CdpServer::new(ServerConfig::default());
460        let (tx, rx) = std::sync::mpsc::channel::<ConsoleMessage>();
461        server.set_console_receiver(rx);
462        assert!(server.console_rx.is_some());
463        // Send a Log message through the channel
464        tx.send(ConsoleMessage::Log {
465            level: "info".into(),
466            text: "hello".into(),
467        })
468        .unwrap();
469        let msg = server.console_rx.as_ref().unwrap().try_recv().unwrap();
470        match msg {
471            ConsoleMessage::Log { level, text } => {
472                assert_eq!(level, "info");
473                assert_eq!(text, "hello");
474            }
475            ConsoleMessage::Event(_) => panic!("expected Log, got Event"),
476        }
477    }
478
479    #[test]
480    fn cdp_server_console_rx_drain_multiple_messages() {
481        let mut server = CdpServer::new(ServerConfig::default());
482        let (tx, rx) = std::sync::mpsc::channel::<ConsoleMessage>();
483        server.set_console_receiver(rx);
484        tx.send(ConsoleMessage::Log {
485            level: "info".into(),
486            text: "msg1".into(),
487        })
488        .unwrap();
489        tx.send(ConsoleMessage::Log {
490            level: "error".into(),
491            text: "msg2".into(),
492        })
493        .unwrap();
494        tx.send(ConsoleMessage::Log {
495            level: "warning".into(),
496            text: "msg3".into(),
497        })
498        .unwrap();
499        let rx_ref = server.console_rx.as_ref().unwrap();
500        let mut messages = Vec::new();
501        while let Ok(msg) = rx_ref.try_recv() {
502            messages.push(msg);
503        }
504        assert_eq!(messages.len(), 3);
505    }
506
507    #[test]
508    fn cdp_server_console_rx_event_variant() {
509        let mut server = CdpServer::new(ServerConfig::default());
510        let (tx, rx) = std::sync::mpsc::channel::<ConsoleMessage>();
511        server.set_console_receiver(rx);
512        tx.send(ConsoleMessage::Event(BaoEvent::PageLoadEventFired {
513            timestamp: 12345.0,
514        }))
515        .unwrap();
516        let msg = server.console_rx.as_ref().unwrap().try_recv().unwrap();
517        match msg {
518            ConsoleMessage::Event(BaoEvent::PageLoadEventFired { timestamp }) => {
519                assert_eq!(timestamp, 12345.0);
520            }
521            other => panic!("expected Event(PageLoadEventFired), got {:?}", other),
522        }
523    }
524
525    #[test]
526    fn cdp_server_console_rx_debugger_script_parsed_event() {
527        let mut server = CdpServer::new(ServerConfig::default());
528        let (tx, rx) = std::sync::mpsc::channel::<ConsoleMessage>();
529        server.set_console_receiver(rx);
530        tx.send(ConsoleMessage::Event(BaoEvent::DebuggerScriptParsed {
531            script_id: "1".into(),
532            url: "test.js".into(),
533            start_line: 0,
534            end_line: 10,
535        }))
536        .unwrap();
537        let msg = server.console_rx.as_ref().unwrap().try_recv().unwrap();
538        match msg {
539            ConsoleMessage::Event(BaoEvent::DebuggerScriptParsed { script_id, url, .. }) => {
540                assert_eq!(script_id, "1");
541                assert_eq!(url, "test.js");
542            }
543            other => panic!("expected Event(DebuggerScriptParsed), got {:?}", other),
544        }
545    }
546
547    #[test]
548    fn cdp_server_console_rx_debugger_paused_event() {
549        let mut server = CdpServer::new(ServerConfig::default());
550        let (tx, rx) = std::sync::mpsc::channel::<ConsoleMessage>();
551        server.set_console_receiver(rx);
552        tx.send(ConsoleMessage::Event(BaoEvent::DebuggerPaused {
553            call_frames: serde_json::json!([]),
554            reason: "breakpoint".into(),
555            hit_breakpoints: serde_json::json!([]),
556        }))
557        .unwrap();
558        let msg = server.console_rx.as_ref().unwrap().try_recv().unwrap();
559        match msg {
560            ConsoleMessage::Event(BaoEvent::DebuggerPaused { reason, .. }) => {
561                assert_eq!(reason, "breakpoint");
562            }
563            other => panic!("expected Event(DebuggerPaused), got {:?}", other),
564        }
565    }
566
567    #[test]
568    fn cdp_server_console_rx_runtime_exception_event() {
569        let mut server = CdpServer::new(ServerConfig::default());
570        let (tx, rx) = std::sync::mpsc::channel::<ConsoleMessage>();
571        server.set_console_receiver(rx);
572        tx.send(ConsoleMessage::Event(BaoEvent::RuntimeExceptionThrown {
573            timestamp: 100.0,
574            text: "TypeError: x is not a function".into(),
575            url: "test.js".into(),
576            line: 10,
577            column: 5,
578            stack_trace: serde_json::Value::Null,
579        }))
580        .unwrap();
581        let msg = server.console_rx.as_ref().unwrap().try_recv().unwrap();
582        match msg {
583            ConsoleMessage::Event(BaoEvent::RuntimeExceptionThrown { text, .. }) => {
584                assert_eq!(text, "TypeError: x is not a function");
585            }
586            other => panic!("expected Event(RuntimeExceptionThrown), got {:?}", other),
587        }
588    }
589
590    #[test]
591    fn cdp_server_console_rx_all_event_variants() {
592        let mut server = CdpServer::new(ServerConfig::default());
593        let (tx, rx) = std::sync::mpsc::channel::<ConsoleMessage>();
594        server.set_console_receiver(rx);
595        let events = vec![
596            ConsoleMessage::Event(BaoEvent::FetchRequestPaused {
597                request_id: "r1".into(),
598                url: "http://test.com".into(),
599                method: "GET".into(),
600                headers: serde_json::json!({}),
601                post_data: None,
602                resource_type: "Document".into(),
603            }),
604            ConsoleMessage::Event(BaoEvent::NetworkRequestWillBeSent {
605                request_id: "req1".into(),
606                url: "http://test.com".into(),
607                method: "GET".into(),
608                headers: serde_json::json!({}),
609                request: serde_json::json!({}),
610                timestamp: 0.0,
611                resource_type: "Document".into(),
612            }),
613            ConsoleMessage::Event(BaoEvent::NetworkResponseReceived {
614                request_id: "req2".into(),
615                url: "http://test.com".into(),
616                status: 200,
617                status_text: "OK".into(),
618                headers: serde_json::json!({}),
619                timestamp: 0.0,
620                resource_type: "Document".into(),
621            }),
622            ConsoleMessage::Event(BaoEvent::NetworkLoadingFailed {
623                request_id: "req3".into(),
624                resource_type: "XHR".into(),
625                error_text: "Network error".into(),
626                timestamp: 0.0,
627            }),
628            ConsoleMessage::Event(BaoEvent::DebuggerScriptParsed {
629                script_id: "1".into(),
630                url: "test.js".into(),
631                start_line: 0,
632                end_line: 10,
633            }),
634            ConsoleMessage::Event(BaoEvent::DebuggerPaused {
635                call_frames: serde_json::json!([]),
636                reason: "other".into(),
637                hit_breakpoints: serde_json::json!([]),
638            }),
639            ConsoleMessage::Event(BaoEvent::RuntimeExceptionThrown {
640                timestamp: 0.0,
641                text: String::new(),
642                url: String::new(),
643                line: 0,
644                column: 0,
645                stack_trace: serde_json::Value::Null,
646            }),
647            ConsoleMessage::Event(BaoEvent::PageLoadEventFired { timestamp: 0.0 }),
648        ];
649        for evt in &events {
650            tx.send(evt.clone()).unwrap();
651        }
652        let rx_ref = server.console_rx.as_ref().unwrap();
653        let mut count = 0;
654        while let Ok(msg) = rx_ref.try_recv() {
655            assert!(matches!(msg, ConsoleMessage::Event(_)));
656            count += 1;
657        }
658        assert_eq!(count, 8);
659    }
660}