1use 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
31const 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
41pub 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 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 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 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 let (mut sink, mut stream) = socket.split();
185 let (tx, mut rx) = tokio::sync::mpsc::channel::<String>(OUTBOUND_QUEUE_CAPACITY);
186
187 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 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 for entry in conn.subs.iter() {
230 entry.value().abort.abort();
231 }
232 }
233
234 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 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 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 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 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 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 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 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#[allow(dead_code)]
535type _UnusedSubBroadcast = broadcast::Sender<()>;