implement subscription
This commit is contained in:
@@ -8,7 +8,7 @@
|
||||
|
||||
import Foundation
|
||||
|
||||
class HomeAssistantDevice: Device
|
||||
class HomeAssistantDevice: Device, Hashable
|
||||
{
|
||||
var name: String
|
||||
var serial: String
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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] = []
|
||||
|
||||
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user