Skip to main content

homecore_api/
ws.rs

1//! WebSocket handler — `/api/websocket`. ADR-130 §2.2 P2 command subset.
2//!
3//! Protocol mirrors HA's WS API:
4//!   server → `{"type":"auth_required","ha_version":"<v>"}`
5//!   client → `{"type":"auth","access_token":"<token>"}`
6//!   server → `{"type":"auth_ok","ha_version":"<v>"}`
7//!   client → `{"id":1,"type":"get_states"}`
8//!   server → `{"id":1,"type":"result","success":true,"result":[...]}`
9//!
10//! `ha_version` is the homecore version string — see ADR-130 Q1 for the
11//! companion-app feature-detect concern.
12//!
13//! ## Security (ADR-161)
14//!
15//! The `auth` token is validated against [`crate::tokens::LongLivedTokenStore`]
16//! via `state.tokens().is_valid()` — the *same* store the REST path uses
17//! (`auth::BearerAuth`). A wrong token receives `auth_invalid` and the socket
18//! is closed. (HC-WS-01 closed the prior bypass where any non-empty token was
19//! accepted.) Command replies are transmitted by a dedicated writer task that
20//! drains the response channel onto the socket (HC-WS-02 closed the prior
21//! reply-theater where responses were logged and discarded).
22
23use std::sync::Arc;
24
25use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
26use axum::extract::State;
27use axum::response::IntoResponse;
28use serde::{Deserialize, Serialize};
29use tokio::sync::broadcast;
30
31/// Per-connection outbound queue. A bounded queue prevents a client that
32/// stops reading from turning event fan-out into unbounded process memory.
33const OUTBOUND_QUEUE_CAPACITY: usize = 256;
34use tracing::warn;
35
36use homecore::{Context, ServiceCall, ServiceName, SystemEvent};
37
38use crate::rest::StateView;
39use crate::state::SharedState;
40
41/// WebSocket upgrade entry point. Mounted on `/api/websocket`.
42pub async fn websocket_handler(
43    ws: WebSocketUpgrade,
44    State(state): State<SharedState>,
45) -> impl IntoResponse {
46    ws.on_upgrade(move |socket| handle_socket(socket, state))
47}
48
49async fn handle_socket(mut socket: WebSocket, state: SharedState) {
50    // Phase 1 — auth handshake.
51    let auth_req = serde_json::json!({
52        "type": "auth_required",
53        "ha_version": state.version(),
54    });
55    if socket
56        .send(Message::Text(auth_req.to_string()))
57        .await
58        .is_err()
59    {
60        return;
61    }
62
63    let token = match socket.recv().await {
64        Some(Ok(Message::Text(raw))) => match serde_json::from_str::<AuthMessage>(&raw) {
65            Ok(m) if m.kind == "auth" => m.access_token,
66            _ => {
67                let _ = socket
68                    .send(Message::Text(
69                        serde_json::json!({"type":"auth_invalid","message":"expected auth"})
70                            .to_string(),
71                    ))
72                    .await;
73                return;
74            }
75        },
76        _ => return,
77    };
78
79    // Validate the bearer token against the same store the REST path
80    // uses (`state.tokens().is_valid()` — see `rest.rs` /
81    // `auth::BearerAuth`). Before the HC-WS-01 fix this checked only
82    // `token.trim().is_empty()` and accepted ANY non-empty token even
83    // with a provisioned `HOMECORE_TOKENS` whitelist — a full WS auth
84    // bypass. `is_valid()` rejects the empty token internally and, in
85    // DEV (`allow_any`) mode, still accepts any non-empty bearer (with
86    // a warn) so smoke tests keep working.
87    if !state.tokens().is_valid(&token).await {
88        let _ = socket
89            .send(Message::Text(
90                serde_json::json!({"type":"auth_invalid","message":"invalid token"}).to_string(),
91            ))
92            .await;
93        return;
94    }
95    let auth_ok = serde_json::json!({"type":"auth_ok","ha_version": state.version()});
96    if socket
97        .send(Message::Text(auth_ok.to_string()))
98        .await
99        .is_err()
100    {
101        return;
102    }
103
104    // Phase 2 — command loop.
105    let conn = Connection::new(state.clone());
106    conn.run(socket).await;
107}
108
109#[derive(Deserialize)]
110struct AuthMessage {
111    #[serde(rename = "type")]
112    kind: String,
113    access_token: String,
114}
115
116#[derive(Deserialize)]
117struct WsCommand {
118    id: u64,
119    #[serde(rename = "type")]
120    kind: String,
121    #[serde(default)]
122    event_type: Option<String>,
123    #[serde(default)]
124    subscription: Option<u64>,
125    #[serde(default)]
126    entity_id: Option<String>,
127    #[serde(default)]
128    domain: Option<String>,
129    #[serde(default)]
130    service: Option<String>,
131    #[serde(default)]
132    service_data: Option<serde_json::Value>,
133    #[serde(default)]
134    event_data: Option<serde_json::Value>,
135    #[serde(default)]
136    template: Option<String>,
137}
138
139#[derive(Serialize)]
140struct ResultMessage<'a> {
141    id: u64,
142    #[serde(rename = "type")]
143    kind: &'static str,
144    success: bool,
145    #[serde(skip_serializing_if = "Option::is_none")]
146    result: Option<serde_json::Value>,
147    #[serde(skip_serializing_if = "Option::is_none")]
148    error: Option<ErrorView<'a>>,
149}
150
151#[derive(Serialize)]
152struct ErrorView<'a> {
153    code: &'static str,
154    message: &'a str,
155}
156
157struct Connection {
158    state: SharedState,
159    subs: Arc<dashmap::DashMap<u64, SubscriptionHandle>>,
160}
161
162struct SubscriptionHandle {
163    abort: tokio::task::AbortHandle,
164}
165
166impl Connection {
167    fn new(state: SharedState) -> Self {
168        Self {
169            state,
170            subs: Arc::new(dashmap::DashMap::new()),
171        }
172    }
173
174    async fn run(self, socket: WebSocket) {
175        use futures_util::{SinkExt, StreamExt};
176
177        let conn = Arc::new(self);
178        // Split the socket so a dedicated writer task can drain `rx` onto
179        // the wire while the reader task processes commands concurrently.
180        // Before the HC-WS-02 fix the socket was moved into a recv-only
181        // task and the only `rx` consumer just `debug!`-logged and
182        // DISCARDED every message — so no `result`/`pong`/`event` ever
183        // reached the client. Now `rx` feeds `socket.send`.
184        let (mut sink, mut stream) = socket.split();
185        let (tx, mut rx) = tokio::sync::mpsc::channel::<String>(OUTBOUND_QUEUE_CAPACITY);
186
187        // Writer task: drain replies onto the socket. A `__pong:<n>`
188        // sentinel maps to a binary Pong control frame; everything else
189        // is a JSON text frame.
190        let writer_task = tokio::spawn(async move {
191            while let Some(msg) = rx.recv().await {
192                let send_result = if let Some(n) = msg.strip_prefix("__pong:") {
193                    let len: usize = n.parse().unwrap_or(0);
194                    sink.send(Message::Pong(vec![0u8; len])).await
195                } else {
196                    sink.send(Message::Text(msg)).await
197                };
198                if send_result.is_err() {
199                    break;
200                }
201            }
202        });
203
204        // Reader task: parse and dispatch commands; responses are pushed
205        // into `tx` and transmitted by the writer task above.
206        let reader_tx = tx.clone();
207        {
208            let conn = Arc::clone(&conn);
209            while let Some(frame) = stream.next().await {
210                match frame {
211                    Ok(Message::Text(raw)) => {
212                        let cmd: WsCommand = match serde_json::from_str(&raw) {
213                            Ok(c) => c,
214                            Err(e) => {
215                                warn!("bad ws command: {e}");
216                                continue;
217                            }
218                        };
219                        conn.handle_cmd(cmd, &reader_tx).await;
220                    }
221                    Ok(Message::Ping(p)) => {
222                        let _ = reader_tx.try_send(format!("__pong:{}", p.len()));
223                    }
224                    Ok(Message::Close(_)) | Err(_) => break,
225                    _ => {}
226                }
227            }
228            // Cancel all subscriptions on disconnect.
229            for entry in conn.subs.iter() {
230                entry.value().abort.abort();
231            }
232        }
233
234        // Reader loop ended → drop the senders so the writer task's `rx`
235        // closes and the task exits cleanly.
236        drop(tx);
237        drop(reader_tx);
238        let _ = writer_task.await;
239    }
240
241    async fn handle_cmd(&self, cmd: WsCommand, tx: &tokio::sync::mpsc::Sender<String>) {
242        match cmd.kind.as_str() {
243            "supported_features" => {
244                // HOMECORE currently emits individual messages. Accepting the
245                // negotiation command keeps modern HA clients compatible while
246                // deliberately declining optional coalescing.
247                self.ack(tx, cmd.id, true, None);
248            }
249            "ping" => {
250                let msg = serde_json::json!({"id": cmd.id, "type": "pong"});
251                let _ = tx.try_send(msg.to_string());
252            }
253            "get_states" => {
254                let snapshots = self.state.homecore().states().all();
255                let views: Vec<StateView> =
256                    snapshots.iter().map(|s| StateView::from_state(s)).collect();
257                self.ack(tx, cmd.id, true, Some(serde_json::to_value(views).unwrap()));
258            }
259            "get_config" => {
260                let payload = serde_json::json!({
261                    "location_name": self.state.location_name(),
262                    "version": self.state.version(),
263                    "state": "RUNNING",
264                });
265                self.ack(tx, cmd.id, true, Some(payload));
266            }
267            "get_panels" => {
268                // Panels are frontend integration resources. An empty map is
269                // the valid shape for a headless server.
270                self.ack(tx, cmd.id, true, Some(serde_json::json!({})));
271            }
272            "get_services" => {
273                let services = self.state.homecore().services().registered_services().await;
274                let mut by_domain: std::collections::HashMap<
275                    String,
276                    serde_json::Map<String, serde_json::Value>,
277                > = std::collections::HashMap::new();
278                for s in services {
279                    by_domain
280                        .entry(s.domain)
281                        .or_default()
282                        .insert(s.service, serde_json::json!({}));
283                }
284                let payload = serde_json::to_value(by_domain).unwrap();
285                self.ack(tx, cmd.id, true, Some(payload));
286            }
287            "config/entity_registry/list" | "get_entity_registry" => {
288                let entries = self.state.homecore().entities().all().await;
289                let payload =
290                    serde_json::to_value(entries).unwrap_or_else(|_| serde_json::json!([]));
291                self.ack(tx, cmd.id, true, Some(payload));
292            }
293            "config/device_registry/list" | "get_device_registry" => {
294                let entries = self.state.homecore().devices().all().await;
295                let payload =
296                    serde_json::to_value(entries).unwrap_or_else(|_| serde_json::json!([]));
297                self.ack(tx, cmd.id, true, Some(payload));
298            }
299            "config/area_registry/list" | "get_area_registry" => {
300                // HOMECORE does not yet model named areas. Returning the valid
301                // empty-list shape lets clients distinguish that from an
302                // unsupported command.
303                self.ack(tx, cmd.id, true, Some(serde_json::json!([])));
304            }
305            "call_service" => {
306                let (Some(domain), Some(service)) = (cmd.domain.clone(), cmd.service.clone())
307                else {
308                    self.err(
309                        tx,
310                        cmd.id,
311                        "missing_domain_service",
312                        "domain and service are required",
313                    );
314                    return;
315                };
316                let call = ServiceCall {
317                    name: ServiceName::new(domain.clone(), service.clone()),
318                    data: cmd.service_data.unwrap_or(serde_json::json!({})),
319                    context: Context::new(),
320                };
321                match self.state.homecore().services().call(call).await {
322                    Ok(v) => self.ack(tx, cmd.id, true, Some(v)),
323                    Err(e) => self.err(tx, cmd.id, "service_error", &e.to_string()),
324                }
325            }
326            "fire_event" => {
327                let Some(event_type) = cmd.event_type.clone() else {
328                    self.err(tx, cmd.id, "invalid_format", "event_type is required");
329                    return;
330                };
331                if !crate::rest::is_valid_event_type(&event_type) {
332                    self.err(tx, cmd.id, "invalid_format", "invalid event_type");
333                    return;
334                }
335                let event_data = cmd.event_data.unwrap_or_else(|| serde_json::json!({}));
336                if !event_data.is_object() {
337                    self.err(tx, cmd.id, "invalid_format", "event_data must be an object");
338                    return;
339                }
340                self.state
341                    .homecore()
342                    .bus()
343                    .fire_domain(homecore::DomainEvent::new(
344                        event_type,
345                        event_data,
346                        Context::new(),
347                    ));
348                self.ack(tx, cmd.id, true, None);
349            }
350            "render_template" => {
351                let Some(template) = cmd.template.as_deref() else {
352                    self.err(tx, cmd.id, "invalid_format", "template is required");
353                    return;
354                };
355                let environment = homecore_automation::TemplateEnvironment::new(Arc::new(
356                    self.state.homecore().states().clone(),
357                ));
358                match environment.render(template) {
359                    Ok(rendered) => {
360                        self.ack(tx, cmd.id, true, Some(serde_json::Value::String(rendered)))
361                    }
362                    Err(error) => self.err(tx, cmd.id, "template_error", &error.to_string()),
363                }
364            }
365            "subscribe_events" => {
366                // HA uses the subscribing command ID as the subscription ID
367                // in every emitted event and in `unsubscribe_events`.
368                let sub_id = cmd.id;
369                if self.subs.contains_key(&sub_id) {
370                    self.err(tx, cmd.id, "id_reused", "subscription id is already active");
371                    return;
372                }
373                let filter = cmd.event_type.clone();
374                let tx_clone = tx.clone();
375                let mut domain_rx = self.state.homecore().bus().subscribe_domain();
376                let mut system_rx = self.state.homecore().bus().subscribe_system();
377                let task = tokio::spawn(async move {
378                    loop {
379                        tokio::select! {
380                            evt = system_rx.recv() => match evt {
381                                Ok(SystemEvent::StateChanged(sc)) => {
382                                    if filter.as_deref() == Some("state_changed") || filter.is_none() {
383                                        let payload = serde_json::json!({
384                                            "id": sub_id,
385                                            "type": "event",
386                                            "event": {
387                                                "event_type": "state_changed",
388                                                "data": {
389                                                    "entity_id": sc.entity_id.as_str(),
390                                                    "old_state": sc.old_state.as_ref().map(|s| StateView::from_state(s)),
391                                                    "new_state": sc.new_state.as_ref().map(|s| StateView::from_state(s)),
392                                                },
393                                                "origin": "LOCAL",
394                                                "time_fired": sc.fired_at.to_rfc3339(),
395                                            }
396                                        });
397                                        if tx_clone.try_send(payload.to_string()).is_err() { break; }
398                                    }
399                                }
400                                Ok(SystemEvent::ServiceCalled { domain, service, data, context }) => {
401                                    if filter.as_deref() == Some("call_service") || filter.is_none() {
402                                        let payload = serde_json::json!({
403                                            "id": sub_id,
404                                            "type": "event",
405                                            "event": {
406                                                "event_type": "call_service",
407                                                "data": {
408                                                    "domain": domain,
409                                                    "service": service,
410                                                    "service_data": data,
411                                                },
412                                                "origin": "LOCAL",
413                                                "time_fired": chrono::Utc::now().to_rfc3339(),
414                                                "context": context,
415                                            }
416                                        });
417                                        if tx_clone.try_send(payload.to_string()).is_err() { break; }
418                                    }
419                                }
420                                Ok(_) => {}
421                                // A slow consumer that falls >4,096 events behind
422                                // gets `Lagged(n)`, which is RECOVERABLE: the bus
423                                // doc (`bus.rs` §"Lagged receivers must re-sync")
424                                // and HA's WS contract both keep the subscription
425                                // alive across a lag. The pre-fix `Err(_) => break`
426                                // treated `Lagged` as fatal, silently killing the
427                                // client's event stream on a burst (HC-WS-LAG-01).
428                                // Skip the dropped window and continue; only a
429                                // `Closed` sender ends the task.
430                                Err(broadcast::error::RecvError::Lagged(_)) => continue,
431                                Err(broadcast::error::RecvError::Closed) => break,
432                            },
433                            evt = domain_rx.recv() => match evt {
434                                Ok(de) => {
435                                    if filter.as_deref() == Some(de.event_type.as_str()) || filter.is_none() {
436                                        let payload = serde_json::json!({
437                                            "id": sub_id,
438                                            "type": "event",
439                                            "event": {
440                                                "event_type": de.event_type,
441                                                "data": de.event_data,
442                                                "origin": format!("{:?}", de.origin).to_uppercase(),
443                                                "time_fired": de.fired_at.to_rfc3339(),
444                                                "context": de.context,
445                                            }
446                                        });
447                                        if tx_clone.try_send(payload.to_string()).is_err() { break; }
448                                    }
449                                }
450                                // Same recoverable-lag handling as the system arm
451                                // above (HC-WS-LAG-01): a lagged domain-event
452                                // receiver re-syncs and continues; only `Closed`
453                                // terminates the subscription.
454                                Err(broadcast::error::RecvError::Lagged(_)) => continue,
455                                Err(broadcast::error::RecvError::Closed) => break,
456                            }
457                        }
458                    }
459                });
460                self.subs.insert(
461                    sub_id,
462                    SubscriptionHandle {
463                        abort: task.abort_handle(),
464                    },
465                );
466                self.ack(tx, cmd.id, true, None);
467            }
468            "unsubscribe_events" => {
469                if let Some(sub_id) = cmd.subscription {
470                    if let Some((_, handle)) = self.subs.remove(&sub_id) {
471                        handle.abort.abort();
472                        self.ack(tx, cmd.id, true, None);
473                    } else {
474                        self.err(tx, cmd.id, "not_found", "subscription_id not found");
475                    }
476                } else {
477                    self.err(
478                        tx,
479                        cmd.id,
480                        "missing_subscription",
481                        "subscription is required",
482                    );
483                }
484            }
485            other => {
486                self.err(
487                    tx,
488                    cmd.id,
489                    "unknown_command",
490                    &format!("unknown ws command: {other}"),
491                );
492            }
493        }
494        // entity_id is reserved for future per-entity subscribes
495        let _ = cmd.entity_id;
496    }
497
498    fn ack(
499        &self,
500        tx: &tokio::sync::mpsc::Sender<String>,
501        id: u64,
502        success: bool,
503        result: Option<serde_json::Value>,
504    ) {
505        let msg = ResultMessage {
506            id,
507            kind: "result",
508            success,
509            result,
510            error: None,
511        };
512        let _ = tx.try_send(serde_json::to_string(&msg).unwrap());
513    }
514
515    fn err(
516        &self,
517        tx: &tokio::sync::mpsc::Sender<String>,
518        id: u64,
519        code: &'static str,
520        message: &str,
521    ) {
522        let msg = ResultMessage {
523            id,
524            kind: "result",
525            success: false,
526            result: None,
527            error: Some(ErrorView { code, message }),
528        };
529        let _ = tx.try_send(serde_json::to_string(&msg).unwrap());
530    }
531}
532
533// Suppress unused warnings for placeholder broadcast type
534#[allow(dead_code)]
535type _UnusedSubBroadcast = broadcast::Sender<()>;