implement subscription

This commit is contained in:
2025-06-19 20:43:05 -07:00
parent dba6754c0c
commit 0d418e76c6
6 changed files with 216 additions and 83 deletions

View File

@@ -8,7 +8,7 @@
import Foundation import Foundation
class HomeAssistantDevice: Device class HomeAssistantDevice: Device, Hashable
{ {
var name: String var name: String
var serial: String var serial: String

View File

@@ -12,6 +12,7 @@ import Foundation
class HomeAssistantServer: Server class HomeAssistantServer: Server
{ {
var connectionStatus: ConnectionStatus = .disconnected var connectionStatus: ConnectionStatus = .disconnected
weak var delegate: ServerDelegate?
// Configure: // Configure:
private let authToken: String = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpc3MiOiI2ZTAwODI1N2Y4N2Q0OWIxYWRkN2ExNTBhOWRmNGZiZCIsImlhdCI6MTc1MDMwMDg5NiwiZXhwIjoyMDY1NjYwODk2fQ.KRf9GpRZdpW9Z_vkbL3sl74rgKS7eAyc8kO0a5jLTgg" private let authToken: String = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpc3MiOiI2ZTAwODI1N2Y4N2Q0OWIxYWRkN2ExNTBhOWRmNGZiZCIsImlhdCI6MTc1MDMwMDg5NiwiZXhwIjoyMDY1NjYwODk2fQ.KRf9GpRZdpW9Z_vkbL3sl74rgKS7eAyc8kO0a5jLTgg"
@@ -23,6 +24,7 @@ class HomeAssistantServer: Server
private let liaison: Liaison private let liaison: Liaison
private var connectionContinuations: [CheckedContinuation<Bool, Never>] = [] private var connectionContinuations: [CheckedContinuation<Bool, Never>] = []
private var eventTask: Task<(), any Error>? = nil private var eventTask: Task<(), any Error>? = nil
private var knownDevices: Set<HomeAssistantDevice> = .init()
typealias NoResult = [String: String] typealias NoResult = [String: String]
@@ -70,10 +72,14 @@ class HomeAssistantServer: Server
await waitConnected() await waitConnected()
do { do {
let registryEntries: [RegistryEntry] = try await liaison.sendMessage(RegistryEntriesMessage()) guard case let .result(registryEntries) = try await liaison.sendMessage(RegistryEntriesMessage())
else { fatalError() }
let filteredEntries = registryEntries.filter { $0.labels.contains(Self.filterLabel) } let filteredEntries = registryEntries.filter { $0.labels.contains(Self.filterLabel) }
let states: [EntityState] = try await liaison.sendMessage(GetStatesMessage()) guard case let .result(states) = try await liaison.sendMessage(GetStatesMessage())
else { fatalError() }
let stateMap = states.reduce(into: [String: DeviceState]()) { partialResult, state in let stateMap = states.reduce(into: [String: DeviceState]()) { partialResult, state in
partialResult[state.entityId] = switch state.state { partialResult[state.entityId] = switch state.state {
case .on: .on case .on: .on
@@ -84,10 +90,12 @@ class HomeAssistantServer: Server
let devices = filteredEntries.map { entry in let devices = filteredEntries.map { entry in
HomeAssistantDevice(entity: entry, state: stateMap[entry.entityId] ?? .on) HomeAssistantDevice(entity: entry, state: stateMap[entry.entityId] ?? .on)
.eraseToAnyDevice()
} }
Task { @MainActor in completion(.success(devices)) } self.knownDevices = Set(devices)
let anyDevices = devices.map { $0.eraseToAnyDevice() }
Task { @MainActor in completion(.success(anyDevices)) }
} catch { } catch {
Task { @MainActor in completion(.failure(error)) } Task { @MainActor in completion(.failure(error)) }
} }
@@ -122,7 +130,7 @@ class HomeAssistantServer: Server
domain: domain domain: domain
) )
let _: NoResult = try await liaison.sendMessage(message) try await liaison.sendMessage(message)
Task { @MainActor in completion(nil) } Task { @MainActor in completion(nil) }
} catch { } catch {
Task { @MainActor in completion(error) } Task { @MainActor in completion(error) }
@@ -131,7 +139,9 @@ class HomeAssistantServer: Server
} }
func responsibleForDevice(_ device: AnyDevice) -> Bool { func responsibleForDevice(_ device: AnyDevice) -> Bool {
return true return knownDevices.contains { knownDevice in
knownDevice.serial == device.serial
}
} }
// MARK: - // MARK: -
@@ -163,8 +173,16 @@ class HomeAssistantServer: Server
case .authRequired: case .authRequired:
Task { await authenticate() } Task { await authenticate() }
case .authOK: case .authOK:
Task { didConnectToWebsocket() } Task {
await subscribeToStateEvents()
didConnectToWebsocket()
}
case .event:
if let subEvent = partialEvent.event {
handleSubscriptionEvent(subEvent)
}
default: default:
print("Unhandled event: \(partialEvent)")
break break
} }
@@ -179,6 +197,25 @@ class HomeAssistantServer: Server
} }
} }
private func handleSubscriptionEvent(_ event: SubscriptionEvent) {
guard let device = knownDevices.first(where: { $0.serial == event.data.entityId }) else { return }
device.state = switch event.data.newState.state {
case .on: .on
case .off: .off
case .other(let string): .off
}
delegate?.server(self, deviceChangedState: device.eraseToAnyDevice())
}
private func subscribeToStateEvents() async {
do {
try await liaison.sendMessage(SubscribeEventsMessage())
} catch {
print("Error subscribing to events: \(error)")
}
}
private func didConnectToWebsocket() { private func didConnectToWebsocket() {
self.connectionStatus = .connected self.connectionStatus = .connected
@@ -190,7 +227,7 @@ class HomeAssistantServer: Server
private func authenticate() async { private func authenticate() async {
do { do {
let _: NoResult = try await liaison.sendMessage(AuthMessage(token: authToken)) try await liaison.sendMessage(AuthMessage(token: authToken))
} catch { } catch {
print("Error authenticating: \(error)") print("Error authenticating: \(error)")
} }
@@ -226,7 +263,7 @@ class HomeAssistantServer: Server
} }
@discardableResult @discardableResult
fileprivate func sendMessage<R: Decodable>(_ message: Message) async throws -> R { fileprivate func sendMessage<M: Message>(_ message: M) async throws -> MessageResult<M.Response> {
return try await withCheckedThrowingContinuation { continuation in return try await withCheckedThrowingContinuation { continuation in
Task { Task {
var outgoingMessage = message var outgoingMessage = message
@@ -234,10 +271,12 @@ class HomeAssistantServer: Server
self.pendingMessages[outgoingMessage.id] = { response in self.pendingMessages[outgoingMessage.id] = { response in
do { do {
let decodedResponse = try JSONDecoder().decode(HomeAssistantServer.Event<R>.self, from: response) let decodedResponse = try JSONDecoder().decode(HomeAssistantServer.Event<M.Response>.self, from: response)
if let result = decodedResponse.result { if let result = decodedResponse.result {
continuation.resume(returning: result) continuation.resume(returning: .result(result))
} else if R.self != NoResult.self { } else if let success = decodedResponse.success {
continuation.resume(with: .success(.success(success)))
} else {
let error = decodedResponse.error ?? ErrorResult(code: "unknown", message: "unknown") let error = decodedResponse.error ?? ErrorResult(code: "unknown", message: "unknown")
continuation.resume(throwing: error) continuation.resume(throwing: error)
} }
@@ -286,8 +325,6 @@ class HomeAssistantServer: Server
websocketTask.resume() websocketTask.resume()
// Event Task // Event Task
let decoder = JSONDecoder()
eventTask?.cancel() eventTask?.cancel()
eventTask = Task.detached { [weak self] in eventTask = Task.detached { [weak self] in
guard let self else { return } guard let self else { return }
@@ -296,11 +333,9 @@ class HomeAssistantServer: Server
switch try await websocketTask.receive() { switch try await websocketTask.receive() {
case .string(let string): case .string(let string):
let data = string.data(using: .utf8)! let data = string.data(using: .utf8)!
let partialEvent = try decoder.decode(PartialEvent.self, from: data) await handlePartialEvent(data: data)
await handlePartialEvent(partialEvent, data: data)
case .data(let data): case .data(let data):
let partialEvent = try decoder.decode(PartialEvent.self, from: data) await handlePartialEvent(data: data)
await handlePartialEvent(partialEvent, data: data)
default: default:
break break
} }
@@ -343,7 +378,9 @@ class HomeAssistantServer: Server
} }
} }
private func handlePartialEvent(_ event: PartialEvent, data: Data) { private func handlePartialEvent(data: Data) {
do {
let event = try JSONDecoder().decode(PartialEvent.self, from: data)
if event.type == .result { if event.type == .result {
guard let id = event.id else { return } guard let id = event.id else { return }
if let resultHandler = self.pendingMessages[id] { if let resultHandler = self.pendingMessages[id] {
@@ -354,6 +391,10 @@ class HomeAssistantServer: Server
} else { } else {
yield(.partialEvent(event, data)) yield(.partialEvent(event, data))
} }
} catch {
// we don't want to bubble this up, because we don't know what the server might send us.
print("PartialEvent decoding error: \(error)")
}
} }
private func yield(_ event: Event) { private func yield(_ event: Event) {
@@ -374,10 +415,11 @@ class HomeAssistantServer: Server
// MARK: - Types // MARK: - Types
struct PartialEvent: Codable struct PartialEvent: Decodable
{ {
let id: Int64? let id: Int64?
let type: EventType let type: EventType
let event: SubscriptionEvent?
} }
struct Event<Result: Decodable>: Decodable struct Event<Result: Decodable>: Decodable
@@ -385,14 +427,45 @@ class HomeAssistantServer: Server
let id: Int? let id: Int?
let type: EventType let type: EventType
let result: Result? let result: Result?
let success: Bool?
let error: ErrorResult? let error: ErrorResult?
} }
struct SubscriptionEvent: Decodable
{
let eventType: SubscriptionEventType
let data: SubscriptionEventData
enum SubscriptionEventType: String, Codable
{
case stateChanged = "state_changed"
}
struct SubscriptionEventData: Decodable
{
let entityId: String
let newState: State
enum CodingKeys: String, CodingKey
{
case entityId = "entity_id"
case newState = "new_state"
}
}
enum CodingKeys: String, CodingKey
{
case eventType = "event_type"
case data
}
}
enum EventType: String, Codable enum EventType: String, Codable
{ {
case authRequired = "auth_required" case authRequired = "auth_required"
case authOK = "auth_ok" case authOK = "auth_ok"
case result = "result" case result = "result"
case event = "event"
} }
struct ErrorResult: Codable, Swift.Error struct ErrorResult: Codable, Swift.Error
@@ -408,23 +481,38 @@ class HomeAssistantServer: Server
// MARK: - Result Types // MARK: - Result Types
struct RegistryEntry: Codable }
{
let id: String
let entityId: String
let name: String?
let labels: [String]
enum CodingKeys: String, CodingKey { extension HomeAssistantDevice
case id {
case name fileprivate convenience init(entity: RegistryEntriesMessage.Entry, state: DeviceState) {
case labels var type: DeviceType = .switch
case entityId = "entity_id" if entity.labels.contains("control_panel_light") {
} type = .light
} }
struct EntityState: Decodable self.init(
{ name: entity.name?.replacingOccurrences(of: " Switch", with: "") ?? entity.entityId,
serial: entity.entityId,
type: type,
state: state
)
}
}
typealias MessageID = Int64?
extension MessageID
{
private static var counter: Int64 = 1
static func next() -> Self {
OSAtomicIncrement64(&Self.counter)
return Self.counter
}
}
struct State: Decodable
{
let entityId: String let entityId: String
let state: State let state: State
@@ -460,41 +548,22 @@ class HomeAssistantServer: Server
case entityId = "entity_id" case entityId = "entity_id"
case state case state
} }
}
}
extension HomeAssistantDevice
{
convenience init(entity: HomeAssistantServer.RegistryEntry, state: DeviceState) {
var type: DeviceType = .switch
if entity.labels.contains("control_panel_light") {
type = .light
}
self.init(
name: entity.name?.replacingOccurrences(of: " Switch", with: "") ?? entity.entityId,
serial: entity.entityId,
type: type,
state: state
)
}
}
typealias MessageID = Int64?
extension MessageID
{
private static var counter: Int64 = 1
static func next() -> Self {
OSAtomicIncrement64(&Self.counter)
return Self.counter
}
} }
// MARK: - Message Types // MARK: - Message Types
typealias Null = String?
enum MessageResult<R>
{
case result(R)
case success(Bool)
}
fileprivate protocol Message: Encodable fileprivate protocol Message: Encodable
{ {
associatedtype Response: Decodable
var id: MessageID { get set } var id: MessageID { get set }
var type: String { get } var type: String { get }
@@ -510,12 +579,31 @@ fileprivate struct RegistryEntriesMessage: Message
{ {
var id: MessageID = nil var id: MessageID = nil
let type: String = "config/entity_registry/list" let type: String = "config/entity_registry/list"
typealias Response = [Entry]
struct Entry: Codable
{
let id: String
let entityId: String
let name: String?
let labels: [String]
enum CodingKeys: String, CodingKey {
case id
case name
case labels
case entityId = "entity_id"
}
}
} }
fileprivate struct GetStatesMessage: Message fileprivate struct GetStatesMessage: Message
{ {
var id: MessageID = nil var id: MessageID = nil
let type: String = "get_states" let type: String = "get_states"
typealias Response = [State]
} }
fileprivate struct AuthMessage: Message fileprivate struct AuthMessage: Message
@@ -532,6 +620,30 @@ fileprivate struct AuthMessage: Message
case type case type
case token = "access_token" case token = "access_token"
} }
typealias Response = Null
}
fileprivate struct SubscribeEventsMessage: Message
{
var id: MessageID = nil
let type: String = "subscribe_events"
let eventType: SubscribeEventType = .stateChanged
enum SubscribeEventType: String, Codable
{
case stateChanged = "state_changed"
}
enum CodingKeys: String, CodingKey
{
case id
case type
case eventType = "event_type"
}
typealias Response = Null
} }
fileprivate struct CallServiceMessage: Message fileprivate struct CallServiceMessage: Message
@@ -570,4 +682,6 @@ fileprivate struct CallServiceMessage: Message
case entityID = "entity_id" case entityID = "entity_id"
} }
} }
typealias Response = Null
} }

View File

@@ -17,6 +17,7 @@ public enum HubitatServerError : Error
public class HubitatServer : Server public class HubitatServer : Server
{ {
var connectionStatus: ConnectionStatus = .disconnected var connectionStatus: ConnectionStatus = .disconnected
weak var delegate: ServerDelegate?
public fileprivate(set) var devices: [HubitatDevice] = [] public fileprivate(set) var devices: [HubitatDevice] = []

View File

@@ -31,6 +31,7 @@ class ServerMultiplex
public func addServer(_ server: Server) public func addServer(_ server: Server)
{ {
server.delegate = self
servers.append(server) servers.append(server)
} }
@@ -138,3 +139,13 @@ extension ServerMultiplex
} }
} }
} }
extension ServerMultiplex: ServerDelegate
{
func server(_ server: any Server, deviceChangedState subjectDevice: AnyDevice) {
guard let device = devices.first(where: { $0.serial == subjectDevice.serial }) else { return }
device.state = subjectDevice.state
delegate?.serverMultiplex(self, devicesStateChanged: [device])
}
}

View File

@@ -16,9 +16,15 @@ enum ConnectionStatus
case error case error
} }
protocol Server protocol ServerDelegate: AnyObject
{
func server(_ server: Server, deviceChangedState: AnyDevice)
}
protocol Server: AnyObject
{ {
var connectionStatus: ConnectionStatus { get } var connectionStatus: ConnectionStatus { get }
var delegate: ServerDelegate? { get set }
/// Designated initializer. Takes an API endpoint URL /// Designated initializer. Takes an API endpoint URL
init(_ url: URL) init(_ url: URL)

View File

@@ -18,6 +18,7 @@ class WemoServer : Server
{ {
private var devices: [WemoDevice] = [] private var devices: [WemoDevice] = []
public var connectionStatus: ConnectionStatus = .disconnected public var connectionStatus: ConnectionStatus = .disconnected
public weak var delegate: ServerDelegate?
fileprivate(set) var baseURL: URL fileprivate(set) var baseURL: URL