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
class HomeAssistantDevice: Device
class HomeAssistantDevice: Device, Hashable
{
var name: String
var serial: String

View File

@@ -12,6 +12,7 @@ import Foundation
class HomeAssistantServer: Server
{
var connectionStatus: ConnectionStatus = .disconnected
weak var delegate: ServerDelegate?
// Configure:
private let authToken: String = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpc3MiOiI2ZTAwODI1N2Y4N2Q0OWIxYWRkN2ExNTBhOWRmNGZiZCIsImlhdCI6MTc1MDMwMDg5NiwiZXhwIjoyMDY1NjYwODk2fQ.KRf9GpRZdpW9Z_vkbL3sl74rgKS7eAyc8kO0a5jLTgg"
@@ -23,6 +24,7 @@ class HomeAssistantServer: Server
private let liaison: Liaison
private var connectionContinuations: [CheckedContinuation<Bool, Never>] = []
private var eventTask: Task<(), any Error>? = nil
private var knownDevices: Set<HomeAssistantDevice> = .init()
typealias NoResult = [String: String]
@@ -70,10 +72,14 @@ class HomeAssistantServer: Server
await waitConnected()
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 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
partialResult[state.entityId] = switch state.state {
case .on: .on
@@ -84,10 +90,12 @@ class HomeAssistantServer: Server
let devices = filteredEntries.map { entry in
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 {
Task { @MainActor in completion(.failure(error)) }
}
@@ -122,7 +130,7 @@ class HomeAssistantServer: Server
domain: domain
)
let _: NoResult = try await liaison.sendMessage(message)
try await liaison.sendMessage(message)
Task { @MainActor in completion(nil) }
} catch {
Task { @MainActor in completion(error) }
@@ -131,7 +139,9 @@ class HomeAssistantServer: Server
}
func responsibleForDevice(_ device: AnyDevice) -> Bool {
return true
return knownDevices.contains { knownDevice in
knownDevice.serial == device.serial
}
}
// MARK: -
@@ -163,8 +173,16 @@ class HomeAssistantServer: Server
case .authRequired:
Task { await authenticate() }
case .authOK:
Task { didConnectToWebsocket() }
Task {
await subscribeToStateEvents()
didConnectToWebsocket()
}
case .event:
if let subEvent = partialEvent.event {
handleSubscriptionEvent(subEvent)
}
default:
print("Unhandled event: \(partialEvent)")
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() {
self.connectionStatus = .connected
@@ -190,7 +227,7 @@ class HomeAssistantServer: Server
private func authenticate() async {
do {
let _: NoResult = try await liaison.sendMessage(AuthMessage(token: authToken))
try await liaison.sendMessage(AuthMessage(token: authToken))
} catch {
print("Error authenticating: \(error)")
}
@@ -226,7 +263,7 @@ class HomeAssistantServer: Server
}
@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
Task {
var outgoingMessage = message
@@ -234,10 +271,12 @@ class HomeAssistantServer: Server
self.pendingMessages[outgoingMessage.id] = { response in
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 {
continuation.resume(returning: result)
} else if R.self != NoResult.self {
continuation.resume(returning: .result(result))
} else if let success = decodedResponse.success {
continuation.resume(with: .success(.success(success)))
} else {
let error = decodedResponse.error ?? ErrorResult(code: "unknown", message: "unknown")
continuation.resume(throwing: error)
}
@@ -286,8 +325,6 @@ class HomeAssistantServer: Server
websocketTask.resume()
// Event Task
let decoder = JSONDecoder()
eventTask?.cancel()
eventTask = Task.detached { [weak self] in
guard let self else { return }
@@ -296,11 +333,9 @@ class HomeAssistantServer: Server
switch try await websocketTask.receive() {
case .string(let string):
let data = string.data(using: .utf8)!
let partialEvent = try decoder.decode(PartialEvent.self, from: data)
await handlePartialEvent(partialEvent, data: data)
await handlePartialEvent(data: data)
case .data(let data):
let partialEvent = try decoder.decode(PartialEvent.self, from: data)
await handlePartialEvent(partialEvent, data: data)
await handlePartialEvent(data: data)
default:
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 {
guard let id = event.id else { return }
if let resultHandler = self.pendingMessages[id] {
@@ -354,6 +391,10 @@ class HomeAssistantServer: Server
} else {
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) {
@@ -374,10 +415,11 @@ class HomeAssistantServer: Server
// MARK: - Types
struct PartialEvent: Codable
struct PartialEvent: Decodable
{
let id: Int64?
let type: EventType
let event: SubscriptionEvent?
}
struct Event<Result: Decodable>: Decodable
@@ -385,14 +427,45 @@ class HomeAssistantServer: Server
let id: Int?
let type: EventType
let result: Result?
let success: Bool?
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
{
case authRequired = "auth_required"
case authOK = "auth_ok"
case result = "result"
case event = "event"
}
struct ErrorResult: Codable, Swift.Error
@@ -408,23 +481,38 @@ class HomeAssistantServer: Server
// MARK: - Result Types
struct RegistryEntry: 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"
}
extension HomeAssistantDevice
{
fileprivate convenience init(entity: RegistryEntriesMessage.Entry, state: DeviceState) {
var type: DeviceType = .switch
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 state: State
@@ -460,41 +548,22 @@ class HomeAssistantServer: Server
case entityId = "entity_id"
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
typealias Null = String?
enum MessageResult<R>
{
case result(R)
case success(Bool)
}
fileprivate protocol Message: Encodable
{
associatedtype Response: Decodable
var id: MessageID { get set }
var type: String { get }
@@ -510,12 +579,31 @@ fileprivate struct RegistryEntriesMessage: Message
{
var id: MessageID = nil
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
{
var id: MessageID = nil
let type: String = "get_states"
typealias Response = [State]
}
fileprivate struct AuthMessage: Message
@@ -532,6 +620,30 @@ fileprivate struct AuthMessage: Message
case type
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
@@ -570,4 +682,6 @@ fileprivate struct CallServiceMessage: Message
case entityID = "entity_id"
}
}
typealias Response = Null
}

View File

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

View File

@@ -31,6 +31,7 @@ class ServerMultiplex
public func addServer(_ server: Server)
{
server.delegate = self
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
}
protocol Server
protocol ServerDelegate: AnyObject
{
func server(_ server: Server, deviceChangedState: AnyDevice)
}
protocol Server: AnyObject
{
var connectionStatus: ConnectionStatus { get }
var delegate: ServerDelegate? { get set }
/// Designated initializer. Takes an API endpoint URL
init(_ url: URL)

View File

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