Skip to main content

active_call/handler/
handler.rs

1use crate::{
2    app::AppState,
3    call::{
4        ActiveCall, ActiveCallType, Command,
5        active_call::{ActiveCallGuard, CallParams},
6    },
7    handler::playbook,
8    playbook::{Playbook, PlaybookRunner},
9};
10use crate::{event::SessionEvent, media::track::TrackConfig};
11use axum::{
12    Json, Router,
13    extract::{Path, Query, State, WebSocketUpgrade, ws::Message},
14    response::sse::{Event, KeepAlive, Sse},
15    response::{IntoResponse, Response},
16    routing::{get, post},
17};
18use bytes::Bytes;
19use chrono::Utc;
20use futures::{SinkExt, StreamExt};
21use rustrtc::IceServer;
22use serde_json::json;
23use std::collections::HashMap;
24use std::{path::PathBuf, sync::Arc, time::Duration};
25use tokio::{join, select};
26use tokio_util::sync::CancellationToken;
27use tracing::{debug, info, trace, warn};
28use uuid::Uuid;
29
30fn filter_headers(
31    extras: &mut std::collections::HashMap<String, serde_json::Value>,
32    allowed_headers: &[String],
33) {
34    extras.retain(|k, _| allowed_headers.iter().any(|h| h.eq_ignore_ascii_case(k)));
35}
36
37pub fn call_router() -> Router<AppState> {
38    let r = Router::new()
39        .route("/call", get(ws_handler))
40        .route("/call/webrtc", get(webrtc_handler))
41        .route("/call/sip", get(sip_handler))
42        .route("/list", get(list_active_calls))
43        .route("/kill/{id}", get(kill_active_call))
44        .route("/events/{id}", get(stream_events))
45        .route("/command/{id}", post(send_command));
46    r
47}
48
49pub fn iceservers_router() -> Router<AppState> {
50    let r = Router::new();
51    r.route("/iceservers", get(get_iceservers))
52}
53
54pub fn playbook_router() -> Router<AppState> {
55    Router::new()
56        .route("/api/playbooks", get(playbook::list_playbooks))
57        .route(
58            "/api/playbooks/{name}",
59            get(playbook::get_playbook).post(playbook::save_playbook),
60        )
61        .route(
62            "/api/playbook/run",
63            axum::routing::post(playbook::run_playbook),
64        )
65        .route("/api/records", get(playbook::list_records))
66}
67
68pub async fn ws_handler(
69    ws: WebSocketUpgrade,
70    State(state): State<AppState>,
71    Query(params): Query<CallParams>,
72) -> Response {
73    call_handler(ActiveCallType::WebSocket, ws, state, params).await
74}
75
76pub async fn sip_handler(
77    ws: WebSocketUpgrade,
78    State(state): State<AppState>,
79    Query(params): Query<CallParams>,
80) -> Response {
81    call_handler(ActiveCallType::Sip, ws, state, params).await
82}
83
84pub async fn webrtc_handler(
85    ws: WebSocketUpgrade,
86    State(state): State<AppState>,
87    Query(params): Query<CallParams>,
88) -> Response {
89    call_handler(ActiveCallType::Webrtc, ws, state, params).await
90}
91
92/// Core call handling logic that works with either WebSocket or mpsc channels
93///
94/// `extras` and `playbook_name` are session-scoped parameters passed directly
95/// by the caller (SIP handler, CLI, etc.) instead of through global maps.
96/// Returns the final call extras (including `_hangup_headers` if set) so the
97/// caller can use them for SIP BYE or other post-call processing.
98pub async fn call_handler_core(
99    call_type: ActiveCallType,
100    session_id: String,
101    app_state: AppState,
102    cancel_token: CancellationToken,
103    audio_receiver: tokio::sync::mpsc::UnboundedReceiver<Bytes>,
104    server_side_track: Option<String>,
105    dump_events: bool,
106    ping_interval: u64,
107    mut command_receiver: tokio::sync::mpsc::UnboundedReceiver<Command>,
108    event_sender_to_client: tokio::sync::mpsc::UnboundedSender<crate::event::SessionEvent>,
109    extras: Option<HashMap<String, serde_json::Value>>,
110    playbook_name: Option<String>,
111) -> Option<HashMap<String, serde_json::Value>> {
112    let _cancel_guard = cancel_token.clone().drop_guard();
113    let track_config = TrackConfig::default();
114
115    let active_call = Arc::new(ActiveCall::new(
116        call_type.clone(),
117        cancel_token.clone(),
118        session_id.clone(),
119        app_state.invitation.clone(),
120        app_state.clone(),
121        track_config,
122        Some(audio_receiver),
123        dump_events,
124        server_side_track,
125        extras,
126        None,
127    ));
128
129    // Load playbook: prefer direct parameter, fall back to pending_playbooks
130    // (pending_playbooks is used by the run_playbook HTTP endpoint)
131    {
132        let name_or_content = playbook_name.or_else(|| {
133            app_state
134                .pending_playbooks
135                .try_lock()
136                .ok()
137                .and_then(|mut pending| pending.remove(&session_id).map(|(val, _)| val))
138        });
139        if let Some(name_or_content) = name_or_content {
140            let playbook_result = if name_or_content.trim().starts_with("---") {
141                Playbook::parse(&name_or_content)
142            } else {
143                // If path already contains config/playbook, use it as-is; otherwise prepend it
144                let path = if name_or_content.starts_with("config/playbook/") {
145                    PathBuf::from(&name_or_content)
146                } else {
147                    PathBuf::from("config/playbook").join(&name_or_content)
148                };
149                Playbook::load(path).await
150            };
151
152            match playbook_result {
153                Ok(mut playbook) => {
154                    // Filter extracted headers if configured (only for SIP calls)
155                    if call_type == ActiveCallType::Sip {
156                        if let Some(sip_config) = &playbook.config.sip {
157                            if let Some(allowed_headers) = &sip_config.extract_headers {
158                                let mut state = active_call.call_state.write().await;
159                                if let Some(extras) = &mut state.extras {
160                                    filter_headers(extras, allowed_headers);
161                                    // Store the list of SIP header keys for later template rendering
162                                    let header_keys: Vec<String> = extras
163                                        .keys()
164                                        .filter(|k| !k.starts_with('_'))
165                                        .cloned()
166                                        .collect();
167                                    extras.insert(
168                                        "_sip_header_keys".to_string(),
169                                        serde_json::to_value(&header_keys).unwrap_or_default(),
170                                    );
171                                    if let Ok(result) = playbook.render(extras) {
172                                        playbook = result;
173                                    }
174                                }
175                            }
176                        }
177                    }
178
179                    match PlaybookRunner::new(playbook, active_call.clone()) {
180                        Ok(runner) => {
181                            crate::spawn(async move {
182                                runner.run().await;
183                            });
184                            let display_name = if name_or_content.trim().starts_with("---") {
185                                "custom content"
186                            } else {
187                                &name_or_content
188                            };
189                            info!(session_id, "Playbook runner started for {}", display_name);
190                        }
191                        Err(e) => {
192                            let display_name = if name_or_content.trim().starts_with("---") {
193                                "custom content"
194                            } else {
195                                &name_or_content
196                            };
197                            warn!(
198                                session_id,
199                                "Failed to create runner {}: {}", display_name, e
200                            )
201                        }
202                    }
203                }
204                Err(e) => {
205                    let display_name = if name_or_content.trim().starts_with("---") {
206                        "custom content"
207                    } else {
208                        &name_or_content
209                    };
210                    warn!(
211                        session_id,
212                        "Failed to load playbook {}: {}", display_name, e
213                    );
214                    let event = SessionEvent::Error {
215                        timestamp: crate::media::get_timestamp(),
216                        track_id: session_id.clone(),
217                        sender: "playbook".to_string(),
218                        error: format!("{}", e),
219                        code: None,
220                    };
221                    event_sender_to_client.send(event).ok();
222                    return None;
223                }
224            }
225        }
226    }
227
228    let recv_commands_loop = async {
229        while let Some(command) = command_receiver.recv().await {
230            if let Err(_) = active_call.enqueue_command(command).await {
231                break;
232            }
233        }
234    };
235
236    let mut event_receiver = active_call.event_sender.subscribe();
237    let send_events_loop = async {
238        loop {
239            match event_receiver.recv().await {
240                Ok(event) => {
241                    if let Err(_) = event_sender_to_client.send(event) {
242                        break;
243                    }
244                }
245                Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => continue,
246                Err(_) => break,
247            }
248        }
249    };
250
251    let send_ping_loop = async {
252        if ping_interval == 0 {
253            active_call.cancel_token.cancelled().await;
254            return;
255        }
256        let mut ticker = tokio::time::interval(Duration::from_secs(ping_interval));
257        loop {
258            ticker.tick().await;
259            let payload = Utc::now().to_rfc3339();
260            let event = SessionEvent::Ping {
261                timestamp: crate::media::get_timestamp(),
262                payload: Some(payload),
263            };
264            if let Err(_) = active_call.event_sender.send(event) {
265                break;
266            }
267        }
268    };
269
270    let guard = ActiveCallGuard::new(active_call.clone());
271    info!(
272        session_id,
273        active_calls = guard.active_calls,
274        ?call_type,
275        "new call started"
276    );
277    let receiver = active_call.new_receiver();
278
279    let (r, _) = join! {
280        active_call.serve(receiver),
281        async {
282            select!{
283                _ = send_ping_loop => {},
284                _ = cancel_token.cancelled() => {},
285                _ = send_events_loop => { },
286                _ = recv_commands_loop => {
287                    info!(session_id, "Command receiver closed");
288                },
289            }
290            cancel_token.cancel();
291        }
292    };
293    // drain events
294    while let Ok(event) = event_receiver.try_recv() {
295        if let Err(_) = event_sender_to_client.send(event) {
296            break;
297        }
298    }
299    match r {
300        Ok(_) => info!(session_id, "call ended successfully"),
301        Err(e) => warn!(session_id, "call ended with error: {}", e),
302    }
303
304    // Capture final extras (including _hangup_headers) before cleanup
305    let final_extras = active_call.call_state.read().await.extras.clone();
306
307    active_call.cleanup().await.ok();
308    debug!(session_id, "Call handler core completed");
309
310    final_extras
311}
312
313pub async fn call_handler(
314    call_type: ActiveCallType,
315    ws: WebSocketUpgrade,
316    app_state: AppState,
317    params: CallParams,
318) -> Response {
319    let session_id = params
320        .id
321        .unwrap_or_else(|| format!("s.{}", Uuid::new_v4().to_string()));
322    let server_side_track = params.server_side_track.clone();
323    let dump_events = params.dump_events.unwrap_or(true);
324    let ping_interval = params.ping_interval.unwrap_or(20);
325
326    let resp = ws.on_upgrade(move |socket| async move {
327        let (mut ws_sender, mut ws_receiver) = socket.split();
328        let (audio_sender, audio_receiver) = tokio::sync::mpsc::unbounded_channel::<Bytes>();
329        let (command_sender, command_receiver) = tokio::sync::mpsc::unbounded_channel::<Command>();
330        let (event_sender_to_client, mut event_receiver_from_core) =
331            tokio::sync::mpsc::unbounded_channel::<crate::event::SessionEvent>();
332        let cancel_token = CancellationToken::new();
333
334        // Start core handler in background
335        let session_id_clone = session_id.clone();
336        let app_state_clone = app_state.clone();
337        let cancel_token_clone = cancel_token.clone();
338        crate::spawn(async move {
339            call_handler_core(
340                call_type,
341                session_id_clone,
342                app_state_clone,
343                cancel_token_clone,
344                audio_receiver,
345                server_side_track,
346                dump_events,
347                ping_interval.into(),
348                command_receiver,
349                event_sender_to_client,
350                None, // extras — not used for WebSocket calls
351                None, // playbook_name — falls back to pending_playbooks
352            )
353            .await;
354        });
355
356        // Handle WebSocket I/O
357        let recv_from_ws_loop = async {
358            while let Some(Ok(message)) = ws_receiver.next().await {
359                match message {
360                    Message::Text(text) => {
361                        let command = match serde_json::from_str::<Command>(&text) {
362                            Ok(cmd) => cmd,
363                            Err(e) => {
364                                warn!(session_id, %text, "Failed to parse command {}", e);
365                                continue;
366                            }
367                        };
368                        if let Err(_) = command_sender.send(command) {
369                            break;
370                        }
371                    }
372                    Message::Binary(bin) => {
373                        audio_sender.send(bin.into()).ok();
374                    }
375                    Message::Close(_) => {
376                        info!(session_id, "WebSocket closed by client");
377                        break;
378                    }
379                    _ => {}
380                }
381            }
382        };
383
384        let send_to_ws_loop = async {
385            while let Some(event) = event_receiver_from_core.recv().await {
386                trace!(session_id, %event, "Sending WS message");
387                let message = match event.into_ws_message() {
388                    Ok(msg) => msg,
389                    Err(e) => {
390                        warn!(session_id, error=%e, "Failed to serialize event to WS message");
391                        continue;
392                    }
393                };
394                if let Err(_) = ws_sender.send(message).await {
395                    info!(session_id, "WebSocket send failed, closing");
396                    break;
397                }
398            }
399        };
400
401        select! {
402            _ = recv_from_ws_loop => {
403                info!(session_id, "WebSocket receive loop ended");
404            },
405            _ = send_to_ws_loop => {
406                info!(session_id, "WebSocket send loop ended");
407            },
408        }
409
410        cancel_token.cancel();
411        ws_sender.flush().await.ok();
412        ws_sender.close().await.ok();
413        debug!(session_id, "WebSocket connection closed");
414    });
415    resp
416}
417
418pub(crate) async fn get_iceservers(State(state): State<AppState>) -> Response {
419    if let Some(ice_servers) = state.config.ice_servers.as_ref() {
420        return Json(ice_servers).into_response();
421    }
422    Json(vec![IceServer {
423        urls: vec!["stun:stun.l.google.com:19302".to_string()],
424        ..Default::default()
425    }])
426    .into_response()
427}
428
429pub(crate) async fn list_active_calls(State(state): State<AppState>) -> Response {
430    let calls = state
431        .active_calls
432        .lock()
433        .unwrap()
434        .iter()
435        .map(|(_, c)| {
436            if let Ok(cs) = c.call_state.try_read() {
437                json!({
438                    "id": c.session_id,
439                    "callType": c.call_type,
440                    "cs.option": cs.option,
441                    "ringTime": cs.ring_time,
442                    "startTime": cs.answer_time,
443                })
444            } else {
445                json!({
446                    "id": c.session_id,
447                    "callType": c.call_type,
448                    "status": "locked",
449                })
450            }
451        })
452        .collect::<Vec<_>>();
453    Json(serde_json::json!({ "active_calls": calls })).into_response()
454}
455
456pub(crate) async fn kill_active_call(
457    Path(id): Path<String>,
458    State(state): State<AppState>,
459) -> Response {
460    let active_calls = state.active_calls.lock().unwrap();
461    if let Some(call) = active_calls.get(&id) {
462        call.cancel_token.cancel();
463        Json(serde_json::json!({ "status": "killed", "id": id })).into_response()
464    } else {
465        (
466            axum::http::StatusCode::NOT_FOUND,
467            Json(serde_json::json!({ "status": "not_found", "id": id })),
468        )
469            .into_response()
470    }
471}
472
473pub(crate) async fn stream_events(
474    Path(id): Path<String>,
475    State(state): State<AppState>,
476) -> Response {
477    let mut rx_events;
478    let mut rx_commands;
479    {
480        let active_calls = state.active_calls.lock().unwrap();
481        if let Some(call) = active_calls.get(&id) {
482            rx_events = call.event_sender.subscribe();
483            rx_commands = call.cmd_sender.subscribe();
484        } else {
485            return (axum::http::StatusCode::NOT_FOUND, "track not active").into_response();
486        }
487    }
488
489    let stream = async_stream::stream! {
490        loop {
491            let result = tokio::select! {
492                r = rx_events.recv() => r.map(|e| serde_json::to_string(&e).map(|json| Event::default().event("event").data(json))),
493                r = rx_commands.recv() => r.map(|c| serde_json::to_string(&c).map(|json| Event::default().event("command").data(json))),
494            };
495            match result {
496                Ok(Ok(sse_event)) => yield Ok::<Event, serde_json::Error>(sse_event),
497                Ok(Err(e)) => yield Err(e.into()),
498                Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
499                Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => continue,
500            }
501        }
502    };
503
504    let mut response = Sse::new(stream)
505        .keep_alive(KeepAlive::default())
506        .into_response();
507    response.headers_mut().insert(
508        axum::http::header::CONTENT_TYPE,
509        "text/event-stream;charset=utf-8".parse().unwrap(),
510    );
511    response
512}
513
514pub(crate) async fn send_command(
515    Path(id): Path<String>,
516    State(state): State<AppState>,
517    Json(command): Json<Command>,
518) -> Response {
519    let active_calls = state.active_calls.lock().unwrap();
520    if let Some(call) = active_calls.get(&id) {
521        if let Ok(_) = call.cmd_sender.send(command) {
522            return Json(serde_json::json!({ "status": "sent", "id": id })).into_response();
523        }
524    }
525
526    (
527        axum::http::StatusCode::NOT_FOUND,
528        Json(serde_json::json!({ "status": "not_found", "id": id })),
529    )
530        .into_response()
531}
532
533trait IntoWsMessage {
534    fn into_ws_message(self) -> Result<Message, serde_json::Error>;
535}
536
537impl IntoWsMessage for crate::event::SessionEvent {
538    fn into_ws_message(self) -> Result<Message, serde_json::Error> {
539        match self {
540            SessionEvent::Binary { data, .. } => Ok(Message::Binary(data.into())),
541            SessionEvent::Ping { timestamp, payload } => {
542                let payload = payload.unwrap_or_else(|| timestamp.to_string());
543                Ok(Message::Ping(payload.into()))
544            }
545            event => serde_json::to_string(&event).map(|payload| Message::Text(payload.into())),
546        }
547    }
548}
549
550#[cfg(test)]
551mod tests {
552    use super::*;
553    use serde_json::json;
554    use std::collections::HashMap;
555
556    #[test]
557    fn test_filter_headers() {
558        let mut extras = HashMap::new();
559        extras.insert("X-Tenant-ID".to_string(), json!("123"));
560        extras.insert("X-User-ID".to_string(), json!("456"));
561        extras.insert("Custom-Header".to_string(), json!("abc"));
562        extras.insert("Irrelevant-Header".to_string(), json!("ignore"));
563
564        // Test case-insensitive matching
565        let allowed = vec!["x-tenant-id".to_string(), "Custom-Header".to_string()];
566
567        filter_headers(&mut extras, &allowed);
568
569        assert!(extras.contains_key("X-Tenant-ID"));
570        assert!(extras.contains_key("Custom-Header"));
571        assert!(!extras.contains_key("X-User-ID"));
572        assert!(!extras.contains_key("Irrelevant-Header"));
573
574        // ensure values are preserved
575        assert_eq!(extras.get("X-Tenant-ID").unwrap(), &json!("123"));
576        assert_eq!(extras.get("Custom-Header").unwrap(), &json!("abc"));
577    }
578
579    #[tokio::test]
580    async fn test_call_handler_core_extras_are_session_scoped() {
581        use crate::app::AppStateBuilder;
582        use crate::call::{ActiveCallType, Command};
583        use crate::config::Config;
584
585        let mut config = Config::default();
586        config.udp_port = 0;
587        let app_state = AppStateBuilder::new()
588            .with_config(config)
589            .build()
590            .await
591            .expect("Failed to build app state");
592
593        let session_id = "test-session-scoped".to_string();
594        let cancel_token = CancellationToken::new();
595
596        // Pass extras directly as a parameter (not via global map)
597        let mut extras = HashMap::new();
598        extras.insert("X-Custom".to_string(), json!("value"));
599
600        let (_audio_sender, audio_receiver) = tokio::sync::mpsc::unbounded_channel::<Bytes>();
601        let (command_sender, command_receiver) = tokio::sync::mpsc::unbounded_channel::<Command>();
602        let (event_sender, _event_receiver) =
603            tokio::sync::mpsc::unbounded_channel::<crate::event::SessionEvent>();
604
605        // Send a Hangup command immediately to end the call
606        command_sender
607            .send(Command::Hangup {
608                reason: None,
609                initiator: None,
610                headers: None,
611                refer: None,
612            })
613            .ok();
614        drop(command_sender);
615
616        // Run call_handler_core with extras passed directly
617        let final_extras = call_handler_core(
618            ActiveCallType::Sip,
619            session_id.clone(),
620            app_state.clone(),
621            cancel_token,
622            audio_receiver,
623            None,
624            false,
625            0,
626            command_receiver,
627            event_sender,
628            Some(extras), // extras passed directly
629            None,         // no playbook
630        )
631        .await;
632
633        // Verify that final extras are returned and contain our custom header
634        assert!(final_extras.is_some(), "final extras should be returned");
635        let extras = final_extras.unwrap();
636        assert_eq!(
637            extras.get("X-Custom"),
638            Some(&json!("value")),
639            "session-scoped extras should be preserved"
640        );
641    }
642}