Init
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
"""Shared helpers for the smartWave project."""
|
||||
|
||||
from .db import DRIVER_NAME, Database, connect, execute, fetchall, fetchone
|
||||
from .mqtt import BACKEND_NAME as MQTT_BACKEND_NAME, BrokerClient, connect as mqtt_connect, publish as mqtt_publish
|
||||
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
+102
@@ -0,0 +1,102 @@
|
||||
"""Small sqlite compatibility layer for CPython and MicroPython.
|
||||
|
||||
The module exposes a tiny wrapper around the available sqlite driver so the
|
||||
same code can run on CPython (`sqlite3`) and MicroPython (`sqlite3` or
|
||||
`usqlite`, depending on the port).
|
||||
"""
|
||||
|
||||
try:
|
||||
import sqlite3 as _sqlite
|
||||
DRIVER_NAME = "sqlite3"
|
||||
except ImportError:
|
||||
try:
|
||||
import usqlite as _sqlite
|
||||
DRIVER_NAME = "usqlite"
|
||||
except ImportError as exc:
|
||||
raise ImportError("No sqlite driver found. Expected sqlite3 or usqlite.") from exc
|
||||
|
||||
|
||||
def _connect(database_path, **connect_kwargs):
|
||||
if connect_kwargs:
|
||||
try:
|
||||
return _sqlite.connect(database_path, **connect_kwargs)
|
||||
except TypeError:
|
||||
pass
|
||||
return _sqlite.connect(database_path)
|
||||
|
||||
|
||||
class Database:
|
||||
"""Lightweight connection wrapper with a consistent API."""
|
||||
|
||||
def __init__(self, database_path, **connect_kwargs):
|
||||
self._database_path = database_path
|
||||
self._connect_kwargs = connect_kwargs
|
||||
self._connection = None
|
||||
|
||||
def open(self):
|
||||
if self._connection is None:
|
||||
self._connection = _connect(self._database_path, **self._connect_kwargs)
|
||||
return self._connection
|
||||
|
||||
def close(self):
|
||||
if self._connection is not None:
|
||||
self._connection.close()
|
||||
self._connection = None
|
||||
|
||||
def commit(self):
|
||||
connection = self.open()
|
||||
if hasattr(connection, "commit"):
|
||||
connection.commit()
|
||||
|
||||
def cursor(self):
|
||||
return self.open().cursor()
|
||||
|
||||
def execute(self, sql, params=None):
|
||||
cursor = self.cursor()
|
||||
if params is None:
|
||||
cursor.execute(sql)
|
||||
else:
|
||||
cursor.execute(sql, params)
|
||||
return cursor
|
||||
|
||||
def executemany(self, sql, params_list):
|
||||
cursor = self.cursor()
|
||||
cursor.executemany(sql, params_list)
|
||||
return cursor
|
||||
|
||||
def fetchone(self, sql, params=None):
|
||||
return self.execute(sql, params).fetchone()
|
||||
|
||||
def fetchall(self, sql, params=None):
|
||||
return self.execute(sql, params).fetchall()
|
||||
|
||||
def executescript(self, script):
|
||||
connection = self.open()
|
||||
if hasattr(connection, "executescript"):
|
||||
return connection.executescript(script)
|
||||
raise NotImplementedError("executescript is not available on this sqlite backend")
|
||||
|
||||
def __enter__(self):
|
||||
self.open()
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, traceback):
|
||||
if exc_type is None:
|
||||
self.commit()
|
||||
self.close()
|
||||
|
||||
|
||||
def connect(database_path, **connect_kwargs):
|
||||
return Database(database_path, **connect_kwargs)
|
||||
|
||||
|
||||
def execute(database_path, sql, params=None, **connect_kwargs):
|
||||
return connect(database_path, **connect_kwargs).execute(sql, params)
|
||||
|
||||
|
||||
def fetchone(database_path, sql, params=None, **connect_kwargs):
|
||||
return connect(database_path, **connect_kwargs).fetchone(sql, params)
|
||||
|
||||
|
||||
def fetchall(database_path, sql, params=None, **connect_kwargs):
|
||||
return connect(database_path, **connect_kwargs).fetchall(sql, params)
|
||||
+218
@@ -0,0 +1,218 @@
|
||||
"""Small MQTT compatibility layer for CPython and MicroPython.
|
||||
|
||||
The wrapper keeps the broker host explicit so embedded clients can point to a
|
||||
real IP address instead of localhost.
|
||||
"""
|
||||
|
||||
try:
|
||||
import paho.mqtt.client as _mqtt
|
||||
BACKEND_NAME = "paho"
|
||||
IS_MICROPYTHON = False
|
||||
except ImportError:
|
||||
try:
|
||||
from umqtt.simple import MQTTClient as _MQTTClient
|
||||
BACKEND_NAME = "umqtt.simple"
|
||||
IS_MICROPYTHON = True
|
||||
except ImportError:
|
||||
try:
|
||||
from umqtt.robust import MQTTClient as _MQTTClient
|
||||
BACKEND_NAME = "umqtt.robust"
|
||||
IS_MICROPYTHON = True
|
||||
except ImportError as exc:
|
||||
raise ImportError("No MQTT client found. Expected paho.mqtt or umqtt.") from exc
|
||||
|
||||
|
||||
DEFAULT_PORT = 8884
|
||||
DEFAULT_CA_FILE = "orchestrateur/mqtt/certs/ca.crt"
|
||||
|
||||
|
||||
def _resolve_host(host):
|
||||
if host:
|
||||
return host
|
||||
|
||||
try:
|
||||
import os
|
||||
|
||||
getenv = getattr(os, "getenv", None)
|
||||
if getenv is not None:
|
||||
host = getenv("MQTT_BROKER_HOST")
|
||||
except Exception:
|
||||
host = None
|
||||
|
||||
if not host:
|
||||
raise ValueError("MQTT broker host is required. Pass the broker IP address instead of localhost.")
|
||||
|
||||
return host
|
||||
|
||||
|
||||
def _ensure_bytes(payload):
|
||||
if payload is None:
|
||||
return b""
|
||||
if isinstance(payload, bytes):
|
||||
return payload
|
||||
if isinstance(payload, bytearray):
|
||||
return bytes(payload)
|
||||
return str(payload).encode()
|
||||
|
||||
|
||||
def _read_file_bytes(path):
|
||||
with open(path, "rb") as handle:
|
||||
return handle.read()
|
||||
|
||||
|
||||
class BrokerClient:
|
||||
"""Small MQTT client with a normalised API across runtimes."""
|
||||
|
||||
def __init__(self, host=None, port=DEFAULT_PORT, client_id=None, use_tls=True, cafile=None, certfile=None, keyfile=None, ssl_params=None, tls_insecure=False, username=None, password=None, keepalive=60):
|
||||
self.host = _resolve_host(host)
|
||||
self.port = port
|
||||
self.client_id = client_id
|
||||
self.use_tls = use_tls
|
||||
self.cafile = cafile or DEFAULT_CA_FILE
|
||||
self.certfile = certfile
|
||||
self.keyfile = keyfile
|
||||
self.ssl_params = ssl_params
|
||||
self.tls_insecure = tls_insecure
|
||||
self.username = username
|
||||
self.password = password
|
||||
self.keepalive = keepalive
|
||||
self._client = None
|
||||
self._callback = None
|
||||
self._messages = []
|
||||
|
||||
def set_callback(self, callback):
|
||||
self._callback = callback
|
||||
if self._client is not None and not IS_MICROPYTHON:
|
||||
self._client.on_message = self._on_message
|
||||
|
||||
def _store_message(self, topic, payload, qos=None, retain=False):
|
||||
message = {
|
||||
"topic": topic,
|
||||
"payload": payload,
|
||||
"qos": qos,
|
||||
"retain": retain,
|
||||
}
|
||||
self._messages.append(message)
|
||||
if self._callback is not None:
|
||||
self._callback(message)
|
||||
|
||||
def _on_message(self, client, userdata, msg):
|
||||
self._store_message(msg.topic, msg.payload, getattr(msg, "qos", None), getattr(msg, "retain", False))
|
||||
|
||||
def open(self):
|
||||
if self._client is not None:
|
||||
return self._client
|
||||
|
||||
if IS_MICROPYTHON:
|
||||
ssl_params = self.ssl_params
|
||||
if self.use_tls and ssl_params is None and self.cafile is not None:
|
||||
ssl_params = {"cadata": _read_file_bytes(self.cafile)}
|
||||
client = _MQTTClient(
|
||||
self.client_id or "smartWave-client",
|
||||
self.host,
|
||||
port=self.port,
|
||||
user=self.username,
|
||||
password=self.password,
|
||||
keepalive=self.keepalive,
|
||||
ssl=self.use_tls or ssl_params is not None,
|
||||
ssl_params=ssl_params,
|
||||
)
|
||||
self._client = client
|
||||
return self._client
|
||||
|
||||
client = _mqtt.Client(client_id=self.client_id or "", clean_session=True, protocol=4, transport="tcp")
|
||||
if self.username is not None or self.password is not None:
|
||||
client.username_pw_set(self.username, self.password)
|
||||
if self.use_tls:
|
||||
tls_kwargs = {}
|
||||
if self.cafile is not None:
|
||||
tls_kwargs["ca_certs"] = self.cafile
|
||||
if self.certfile is not None:
|
||||
tls_kwargs["certfile"] = self.certfile
|
||||
if self.keyfile is not None:
|
||||
tls_kwargs["keyfile"] = self.keyfile
|
||||
if tls_kwargs:
|
||||
client.tls_set(**tls_kwargs)
|
||||
else:
|
||||
client.tls_set()
|
||||
if self.tls_insecure:
|
||||
client.tls_insecure_set(True)
|
||||
client.on_message = self._on_message
|
||||
self._client = client
|
||||
return self._client
|
||||
|
||||
def connect(self):
|
||||
client = self.open()
|
||||
if IS_MICROPYTHON:
|
||||
client.connect()
|
||||
return client
|
||||
|
||||
client.connect(self.host, self.port, self.keepalive)
|
||||
return client
|
||||
|
||||
def publish(self, topic, payload, qos=2, retain=False):
|
||||
client = self.open()
|
||||
payload_bytes = _ensure_bytes(payload)
|
||||
if IS_MICROPYTHON:
|
||||
return client.publish(topic, payload_bytes, retain=retain, qos=qos)
|
||||
return client.publish(topic, payload_bytes, qos=qos, retain=retain)
|
||||
|
||||
def subscribe(self, topic, qos=2):
|
||||
client = self.open()
|
||||
if IS_MICROPYTHON:
|
||||
client.set_callback(self._on_micropython_message)
|
||||
return client.subscribe(topic, qos=qos)
|
||||
return client.subscribe(topic, qos=qos)
|
||||
|
||||
def _on_micropython_message(self, topic, payload):
|
||||
self._store_message(topic, payload, None, False)
|
||||
|
||||
def poll(self, timeout=0.1):
|
||||
if self._client is None:
|
||||
return None
|
||||
if IS_MICROPYTHON:
|
||||
return self._client.check_msg()
|
||||
return self._client.loop(timeout=timeout)
|
||||
|
||||
def wait(self):
|
||||
if self._client is None:
|
||||
return None
|
||||
if IS_MICROPYTHON:
|
||||
return self._client.wait_msg()
|
||||
return self._client.loop_forever()
|
||||
|
||||
def get_message(self):
|
||||
if not self._messages:
|
||||
return None
|
||||
return self._messages.pop(0)
|
||||
|
||||
def close(self):
|
||||
if self._client is None:
|
||||
return
|
||||
try:
|
||||
self._client.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
self._client = None
|
||||
|
||||
def __enter__(self):
|
||||
self.connect()
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, traceback):
|
||||
self.close()
|
||||
|
||||
|
||||
def connect(host=None, **client_kwargs):
|
||||
return BrokerClient(host=host, **client_kwargs)
|
||||
|
||||
|
||||
def publish(host, topic, payload, **client_kwargs):
|
||||
qos = client_kwargs.pop("qos", 2)
|
||||
retain = client_kwargs.pop("retain", False)
|
||||
client = connect(host=host, **client_kwargs)
|
||||
client.connect()
|
||||
try:
|
||||
return client.publish(topic, payload, qos=qos, retain=retain)
|
||||
finally:
|
||||
client.close()
|
||||
Reference in New Issue
Block a user