diff --git a/bumper/__init__.py b/bumper/__init__.py index 5b0f8ad..66c9d26 100644 --- a/bumper/__init__.py +++ b/bumper/__init__.py @@ -1,3 +1,4 @@ +"""Init module.""" import asyncio import logging import os @@ -5,8 +6,8 @@ import socket import sys from bumper.db import ( - bot_reset_connectionStatus, - client_reset_connectionStatus, + bot_reset_connection_status, + client_reset_connection_status, revoke_expired_oauths, revoke_expired_tokens, ) @@ -18,6 +19,7 @@ from bumper.xmppserver import XMPPServer def strtobool(strbool: str | bool | None) -> bool: + """Convert str to bool.""" if str(strbool).lower() in ["true", "1", "t", "y", "on", "yes"]: return True else: @@ -77,13 +79,14 @@ web_server_bindings = [ async def start() -> None: + """Start bumper.""" # Reset xmpp/mqtt to false in database for bots and clients - bot_reset_connectionStatus() - client_reset_connectionStatus() + bot_reset_connection_status() + client_reset_connection_status() try: loop = asyncio.get_event_loop() - except: + except: # noqa: E722 loop = asyncio.new_event_loop() if bumper_debug: @@ -150,22 +153,28 @@ async def start() -> None: async def maintenance() -> None: + """Run maintenance.""" revoke_expired_tokens() revoke_expired_oauths() async def shutdown() -> None: + """Shutdown bumper.""" try: bumperlog.info("Shutting down") global shutting_down shutting_down = True + global mqtt_helperbot await mqtt_helperbot.disconnect() + global web_server await web_server.shutdown() + global mqtt_server while mqtt_server.state == "starting": await asyncio.sleep(0.1) if mqtt_server.state == "started": await mqtt_server.shutdown() + global xmpp_server if xmpp_server.server: if xmpp_server.server.is_serving: xmpp_server.server.close() @@ -177,6 +186,7 @@ async def shutdown() -> None: def main(argv: None | list[str] = None) -> None: + """Start everything.""" import argparse global bumper_debug diff --git a/bumper/db.py b/bumper/db.py index 79d92df..d55bda9 100644 --- a/bumper/db.py +++ b/bumper/db.py @@ -1,3 +1,4 @@ +"""Database module.""" import os from datetime import datetime, timedelta from typing import Any @@ -13,17 +14,17 @@ from .util import get_logger _LOGGER = get_logger("db") -def db_file() -> str: - return os.environ.get("DB_FILE") or os_db_path() +def _db_file() -> str: + return os.environ.get("DB_FILE") or _os_db_path() -def os_db_path() -> str: # createdir=True): +def _os_db_path() -> str: # createdir=True): return os.path.join(bumper.data_dir, "bumper.db") -def db_get() -> TinyDB: +def _db_get() -> TinyDB: # Will create the database if it doesn't exist - db = TinyDB(db_file()) + db = TinyDB(_db_file()) # Will create the tables if they don't exist db.table("users", cache_size=0) @@ -36,29 +37,32 @@ def db_get() -> TinyDB: def user_add(userid: str) -> None: + """Add user.""" newuser = BumperUser() newuser.userid = userid user = user_get(userid) if not user: _LOGGER.info(f"Adding new user with userid: {newuser.userid}") - user_full_upsert(newuser.asdict()) + _user_full_upsert(newuser.asdict()) def user_get(userid: str) -> None | Document: - users = db_get().table("users") + """Get user.""" + users = _db_get().table("users") User = Query() return users.get(User.userid == userid) -def user_by_deviceid(deviceid: str) -> None | Document: - users = db_get().table("users") +def user_by_device_id(deviceid: str) -> None | Document: + """Get user by device id.""" + users = _db_get().table("users") User = Query() return users.get(User.devices.any([deviceid])) -def user_full_upsert(user: dict[str, Any]) -> None: - opendb = db_get() +def _user_full_upsert(user: dict[str, Any]) -> None: + opendb = _db_get() with opendb: users = opendb.table("users") User = Query() @@ -66,21 +70,23 @@ def user_full_upsert(user: dict[str, Any]) -> None: def user_add_device(userid: str, devid: str) -> None: - opendb = db_get() + """Add device to user.""" + opendb = _db_get() with opendb: users = opendb.table("users") User = Query() user = users.get(User.userid == userid) if user: userdevices = list(user["devices"]) - if not devid in userdevices: + if devid not in userdevices: userdevices.append(devid) users.upsert({"devices": userdevices}, User.userid == userid) def user_remove_device(userid: str, devid: str) -> None: - opendb = db_get() + """Remove device from user.""" + opendb = _db_get() with opendb: users = opendb.table("users") User = Query() @@ -94,21 +100,23 @@ def user_remove_device(userid: str, devid: str) -> None: def user_add_bot(userid: str, did: str) -> None: - opendb = db_get() + """Add bot to user.""" + opendb = _db_get() with opendb: users = opendb.table("users") User = Query() user = users.get(User.userid == userid) if user: userbots = list(user["bots"]) - if not did in userbots: + if did not in userbots: userbots.append(did) users.upsert({"bots": userbots}, User.userid == userid) def user_remove_bot(userid: str, did: str) -> None: - opendb = db_get() + """Remove bot from user.""" + opendb = _db_get() with opendb: users = opendb.table("users") User = Query() @@ -122,17 +130,20 @@ def user_remove_bot(userid: str, did: str) -> None: def user_get_tokens(userid: str) -> list[Document]: - tokens = db_get().table("tokens") + """Get all tokens by given user.""" + tokens = _db_get().table("tokens") return tokens.search(Query().userid == userid) def user_get_token(userid: str, token: str) -> Document | None: - tokens = db_get().table("tokens") + """Get token by user.""" + tokens = _db_get().table("tokens") return tokens.get((Query().userid == userid) & (Query().token == token)) def user_add_token(userid: str, token: str) -> None: - opendb = db_get() + """Ass token for given user.""" + opendb = _db_get() with opendb: tokens = opendb.table("tokens") tmptoken = tokens.get((Query().userid == userid) & (Query().token == token)) @@ -151,7 +162,8 @@ def user_add_token(userid: str, token: str) -> None: def user_revoke_all_tokens(userid: str) -> None: - opendb = db_get() + """Revoke all tokens for given user.""" + opendb = _db_get() with opendb: tokens = opendb.table("tokens") tsearch = tokens.search(Query().userid == userid) @@ -160,7 +172,8 @@ def user_revoke_all_tokens(userid: str) -> None: def user_revoke_expired_tokens(userid: str) -> None: - opendb = db_get() + """Revoke expired user tokens.""" + opendb = _db_get() with opendb: tokens = opendb.table("tokens") tsearch = tokens.search(Query().userid == userid) @@ -171,7 +184,8 @@ def user_revoke_expired_tokens(userid: str) -> None: def user_revoke_token(userid: str, token: str) -> None: - opendb = db_get() + """Revoke user token.""" + opendb = _db_get() with opendb: tokens = opendb.table("tokens") tmptoken = tokens.get((Query().userid == userid) & (Query().token == token)) @@ -180,7 +194,8 @@ def user_revoke_token(userid: str, token: str) -> None: def user_add_authcode(userid: str, token: str, authcode: str) -> None: - opendb = db_get() + """Add user authcode.""" + opendb = _db_get() with opendb: tokens = opendb.table("tokens") tmptoken = tokens.get((Query().userid == userid) & (Query().token == token)) @@ -192,7 +207,8 @@ def user_add_authcode(userid: str, token: str, authcode: str) -> None: def user_revoke_authcode(userid: str, token: str) -> None: - opendb = db_get() + """Revoke user authcode.""" + opendb = _db_get() with opendb: tokens = opendb.table("tokens") tmptoken = tokens.get((Query().userid == userid) & (Query().token == token)) @@ -204,7 +220,8 @@ def user_revoke_authcode(userid: str, token: str) -> None: def revoke_expired_oauths() -> None: - opendb = db_get() + """Revoke expired oauths.""" + opendb = _db_get() with opendb: table = opendb.table("oauth") entries = table.all() @@ -217,7 +234,8 @@ def revoke_expired_oauths() -> None: def user_revoke_expired_oauths(userid: str) -> None: - opendb = db_get() + """Revoke expired oauths by user.""" + opendb = _db_get() with opendb: table = opendb.table("oauth") search = table.search(Query().userid == userid) @@ -229,8 +247,9 @@ def user_revoke_expired_oauths(userid: str) -> None: def user_add_oauth(userid: str) -> OAuth: + """Add oauth for user.""" user_revoke_expired_oauths(userid) - opendb = db_get() + opendb = _db_get() with opendb: table = opendb.table("oauth") entry = table.get(Query().userid == userid) @@ -244,19 +263,22 @@ def user_add_oauth(userid: str) -> OAuth: def token_by_authcode(authcode: str) -> Document | None: - tokens = db_get().table("tokens") + """Get token by authcode.""" + tokens = _db_get().table("tokens") return tokens.get(Query().authcode == authcode) def get_disconnected_xmpp_clients() -> list[Document]: - clients = db_get().table("clients") - Client = Query() - return clients.search(Client.xmpp_connection == False) + """Get disconnected XMPP clients.""" + clients = _db_get().table("clients") + client = Query() + return clients.search(client.xmpp_connection == False) # noqa: E712 def check_authcode(uid: str, authcode: str) -> bool: + """Check authcode.""" _LOGGER.debug(f"Checking for authcode: {authcode}") - tokens = db_get().table("tokens") + tokens = _db_get().table("tokens") tmpauth = tokens.get( (Query().authcode == authcode) & ( # Match authcode @@ -270,9 +292,10 @@ def check_authcode(uid: str, authcode: str) -> bool: return False -def loginByItToken(authcode: str) -> dict[str, str]: +def login_by_it_token(authcode: str) -> dict[str, str]: + """Login by token.""" _LOGGER.debug(f"Checking for authcode: {authcode}") - tokens = db_get().table("tokens") + tokens = _db_get().table("tokens") tmpauth = tokens.get( Query().authcode == authcode @@ -288,8 +311,9 @@ def loginByItToken(authcode: str) -> dict[str, str]: def check_token(uid: str, token: str) -> bool: + """Check token.""" _LOGGER.debug(f"Checking for token: {token}") - tokens = db_get().table("tokens") + tokens = _db_get().table("tokens") tmpauth = tokens.get( (Query().token == token) & ( # Match token @@ -304,122 +328,137 @@ def check_token(uid: str, token: str) -> bool: def revoke_expired_tokens() -> None: - tokens = db_get().table("tokens").all() + """Revoke expired tokens.""" + tokens = _db_get().table("tokens").all() for i in tokens: if datetime.now() >= datetime.fromisoformat(i["expiration"]): _LOGGER.debug("Removing token {} due to expiration".format(i["token"])) - db_get().table("tokens").remove(doc_ids=[i.doc_id]) + _db_get().table("tokens").remove(doc_ids=[i.doc_id]) -def bot_add(sn: str, did: str, devclass: str, resource: str, company: str) -> None: - newbot = VacBotDevice() - newbot.did = did - newbot.name = sn - newbot.vac_bot_device_class = devclass - newbot.resource = resource - newbot.company = company +def bot_add(sn: str, did: str, dev_class: str, resource: str, company: str) -> None: + """Add bot.""" + new_bot = VacBotDevice() + new_bot.did = did + new_bot.name = sn + new_bot.vac_bot_device_class = dev_class + new_bot.resource = resource + new_bot.company = company bot = bot_get(did) if not bot: # Not existing bot in database if ( - not devclass == "" or "@" not in sn or "tmp" not in sn + not dev_class == "" or "@" not in sn or "tmp" not in sn ): # try to prevent bad additions to the bot list - _LOGGER.info(f"Adding new bot with SN: {newbot.name} DID: {newbot.did}") - bot_full_upsert(newbot.asdict()) + _LOGGER.info(f"Adding new bot with SN: {new_bot.name} DID: {new_bot.did}") + bot_full_upsert(new_bot.asdict()) def bot_remove(did: str) -> None: - bots = db_get().table("bots") + """Remove bot.""" + bots = _db_get().table("bots") bot = bot_get(did) if bot: bots.remove(doc_ids=[bot.doc_id]) def bot_get(did: str) -> Document | None: - bots = db_get().table("bots") - Bot = Query() - return bots.get(Bot.did == did) + """Get bot.""" + bots = _db_get().table("bots") + bot = Query() + return bots.get(bot.did == did) def bot_full_upsert(vacbot: dict[str, Any]) -> None: - bots = db_get().table("bots") - Bot = Query() + """Upsert bot.""" + bots = _db_get().table("bots") + bot = Query() if "did" in vacbot: - bots.upsert(vacbot, Bot.did == vacbot["did"]) + bots.upsert(vacbot, bot.did == vacbot["did"]) else: _LOGGER.error(f"No DID in vacbot: {vacbot}") def bot_set_nick(did: str, nick: str) -> None: - bots = db_get().table("bots") - Bot = Query() - bots.upsert({"nick": nick}, Bot.did == did) + """Bot set nickname.""" + bots = _db_get().table("bots") + bot = Query() + bots.upsert({"nick": nick}, bot.did == did) def bot_set_mqtt(did: str, mqtt: bool) -> None: - bots = db_get().table("bots") - Bot = Query() - bots.upsert({"mqtt_connection": mqtt}, Bot.did == did) + """Bot ste MQTT status.""" + bots = _db_get().table("bots") + bot = Query() + bots.upsert({"mqtt_connection": mqtt}, bot.did == did) def bot_set_xmpp(did: str, xmpp: bool) -> None: - bots = db_get().table("bots") - Bot = Query() - bots.upsert({"xmpp_connection": xmpp}, Bot.did == did) + """Bot set XMPP status.""" + bots = _db_get().table("bots") + bot = Query() + bots.upsert({"xmpp_connection": xmpp}, bot.did == did) def client_add(userid: str, realm: str, resource: str) -> None: - newclient = VacBotClient() - newclient.userid = userid - newclient.realm = realm - newclient.resource = resource + """Add client.""" + new_client = VacBotClient() + new_client.userid = userid + new_client.realm = realm + new_client.resource = resource client = client_get(resource) if not client: - _LOGGER.info(f"Adding new client with resource {newclient.resource}") - client_full_upsert(newclient.asdict()) + _LOGGER.info(f"Adding new client with resource {new_client.resource}") + _client_full_upsert(new_client.asdict()) def client_remove(resource: str) -> None: - clients = db_get().table("clients") + """Remove client.""" + clients = _db_get().table("clients") client = client_get(resource) if client: clients.remove(doc_ids=[client.doc_id]) def client_get(resource: str) -> Document | None: - clients = db_get().table("clients") - Client = Query() - return clients.get(Client.resource == resource) + """Get client by resource.""" + clients = _db_get().table("clients") + client = Query() + return clients.get(client.resource == resource) -def client_full_upsert(client: dict[str, Any]) -> None: - clients = db_get().table("clients") - Client = Query() - clients.upsert(client, Client.resource == client["resource"]) +def _client_full_upsert(client: dict[str, Any]) -> None: + clients = _db_get().table("clients") + client_query = Query() + clients.upsert(client, client_query.resource == client["resource"]) def client_set_mqtt(resource: str, mqtt: bool) -> None: - clients = db_get().table("clients") - Client = Query() - clients.upsert({"mqtt_connection": mqtt}, Client.resource == resource) + """Client set MQTT status.""" + clients = _db_get().table("clients") + client = Query() + clients.upsert({"mqtt_connection": mqtt}, client.resource == resource) def client_set_xmpp(resource: str, xmpp: bool) -> None: - clients = db_get().table("clients") - Client = Query() - clients.upsert({"xmpp_connection": xmpp}, Client.resource == resource) + """Client set XMPP status.""" + clients = _db_get().table("clients") + client = Query() + clients.upsert({"xmpp_connection": xmpp}, client.resource == resource) -def bot_reset_connectionStatus() -> None: - bots = db_get().table("bots") +def bot_reset_connection_status() -> None: + """Reset all bot connection status.""" + bots = _db_get().table("bots") for bot in bots: bot_set_mqtt(bot["did"], False) bot_set_xmpp(bot["did"], False) -def client_reset_connectionStatus() -> None: - clients = db_get().table("clients") +def client_reset_connection_status() -> None: + """Reset all client connection status.""" + clients = _db_get().table("clients") for client in clients: client_set_mqtt(client["resource"], False) client_set_xmpp(client["resource"], False) diff --git a/bumper/dns.py b/bumper/dns.py index c1c071a..acfa52a 100644 --- a/bumper/dns.py +++ b/bumper/dns.py @@ -1,11 +1,14 @@ +"""Dns module.""" from aiohttp import AsyncResolver def get_resolver_with_public_nameserver() -> AsyncResolver: + """Get resolver.""" # requires aiodns return AsyncResolver(nameservers=["1.1.1.1", "8.8.8.8"]) async def resolve(host: str) -> str: + """Resolve host.""" hosts = await get_resolver_with_public_nameserver().resolve(host) return hosts[0]["host"] # type:ignore[no-any-return] diff --git a/bumper/models.py b/bumper/models.py index d711e9c..ea40790 100644 --- a/bumper/models.py +++ b/bumper/models.py @@ -1,3 +1,4 @@ +"""Models module.""" import json import uuid from datetime import datetime, timedelta @@ -8,6 +9,8 @@ from bumper.util import convert_to_millis class VacBotDevice: + """Vacuum device.""" + def __init__( self, did: str = "", @@ -27,6 +30,7 @@ class VacBotDevice: self.xmpp_connection = False def asdict(self) -> dict[str, str | bool]: + """Convert to dict.""" return { "class": self.vac_bot_device_class, "company": self.company, @@ -40,16 +44,21 @@ class VacBotDevice: class BumperUser: + """Bumper user.""" + def __init__(self, userid: str = ""): self.userid = userid self.devices: list[str] = [] self.bots: list[str] = [] def asdict(self) -> dict[str, Any]: + """Convert to dict.""" return {"userid": self.userid, "devices": self.devices, "bots": self.bots} -class GlobalVacBotDevice(VacBotDevice): # EcoVacs Home +class GlobalVacBotDevice(VacBotDevice): + """Global vacuum device.""" + UILogicId = "" ota = True updateInfo = {"changeLog": "", "needUpdate": False} @@ -58,6 +67,8 @@ class GlobalVacBotDevice(VacBotDevice): # EcoVacs Home class VacBotClient: + """Vacuum client.""" + def __init__(self, userid: str = "", realm: str = "", token: str = ""): self.userid = userid self.realm = realm @@ -66,6 +77,7 @@ class VacBotClient: self.xmpp_connection = False def asdict(self) -> dict[str, Any]: + """Convert to dict.""" return { "userid": self.userid, "realm": self.realm, @@ -76,6 +88,8 @@ class VacBotClient: class EcoVacs_Login: + """Ecovacs login.""" + accessToken = "" country = "" email = "" @@ -83,16 +97,21 @@ class EcoVacs_Login: username = "" def toJSON(self) -> str: + """Convert to json.""" return json.dumps(self, default=lambda o: o.__dict__, sort_keys=False) class EcoVacsHome_Login(EcoVacs_Login): + """Ecovacs home login.""" + loginName = "" mobile: str | None = "" ucUid = "" class OAuth: + """Oauth.""" + access_token = "" expire_at = "" refresh_token = "" @@ -103,6 +122,7 @@ class OAuth: @classmethod def create_new(cls, userId: str) -> "OAuth": + """Create new.""" oauth = OAuth() oauth.userId = userId oauth.access_token = uuid.uuid4().hex @@ -113,9 +133,11 @@ class OAuth: return oauth def toDB(self) -> dict: + """Convert for db.""" return self.__dict__ def toResponse(self) -> dict: + """Convert to response.""" data = self.__dict__ data["expire_at"] = convert_to_millis( datetime.fromisoformat(self.expire_at).timestamp() diff --git a/bumper/mqtt/helper_bot.py b/bumper/mqtt/helper_bot.py index d7d231b..f52fd86 100644 --- a/bumper/mqtt/helper_bot.py +++ b/bumper/mqtt/helper_bot.py @@ -154,6 +154,7 @@ class HelperBot: self._commands.pop(request_id, None) def publish(self, topic: str, data: bytes) -> None: + """Publish message.""" self._client.publish(topic, data) async def disconnect(self) -> None: diff --git a/bumper/mqtt/proxy.py b/bumper/mqtt/proxy.py index aca73bc..d6a856e 100644 --- a/bumper/mqtt/proxy.py +++ b/bumper/mqtt/proxy.py @@ -1,6 +1,5 @@ """Mqtt proxy module.""" import asyncio -import re import ssl import typing from collections.abc import MutableMapping @@ -52,6 +51,7 @@ class ProxyClient: self._port = port async def connect(self, username: str, password: str) -> None: + """Connect.""" try: await self._client.connect( f"mqtts://{username}:{password}@{self._host}:{self._port}" @@ -96,12 +96,15 @@ class ProxyClient: ) async def subscribe(self, topic: str, qos: QOS_0 | QOS_1 | QOS_2 = QOS_0) -> None: + """Subscribe to topic.""" await self._client.subscribe([(topic, qos)]) async def disconnect(self) -> None: + """Disconnect.""" await self._client.disconnect() async def publish(self, topic: str, message: bytes, qos: int | None = None) -> None: + """Publish message.""" await self._client.publish(topic, message, qos) diff --git a/bumper/mqtt/server.py b/bumper/mqtt/server.py index d831bd9..8556454 100644 --- a/bumper/mqtt/server.py +++ b/bumper/mqtt/server.py @@ -253,6 +253,8 @@ class BumperMQTTServerPlugin: async def on_broker_client_subscribed( self, client_id: str, topic: str, qos: QOS_0 | QOS_1 | QOS_2 ) -> None: + """Is called when a client subscribes on the broker.""" + if bumper.bumper_proxy_mqtt: # if proxy mode, also subscribe on ecovacs server if client_id in self._proxy_clients: diff --git a/bumper/util.py b/bumper/util.py index 088945d..3a96509 100644 --- a/bumper/util.py +++ b/bumper/util.py @@ -1,3 +1,5 @@ +"""Util module.""" + import logging import os import sys @@ -14,6 +16,7 @@ log_to_stdout = os.environ.get("LOG_TO_STDOUT") def get_logger(name: str, rotate: RotatingFileHandler | None = None) -> logging.Logger: + """Get logger.""" found_logger = __loggers.get(name) if found_logger: return found_logger @@ -53,4 +56,5 @@ def convert_to_millis(seconds: int | float) -> int: def get_current_time_as_millis() -> int: + """Get current time in millis.""" return convert_to_millis(datetime.utcnow().timestamp()) diff --git a/bumper/web/auth_util.py b/bumper/web/auth_util.py index 3352f55..4b99c64 100644 --- a/bumper/web/auth_util.py +++ b/bumper/web/auth_util.py @@ -11,13 +11,13 @@ from aiohttp.web_response import Response from bumper import db, use_auth from bumper.db import ( - db_get, + _db_get, user_add, user_add_authcode, user_add_bot, user_add_device, user_add_token, - user_by_deviceid, + user_by_device_id, user_get, user_get_token, user_revoke_expired_tokens, @@ -60,7 +60,7 @@ async def login(request: Request) -> Response: if ( not user_devid == "" ): # Performing basic "auth" using devid, super insecure - user = user_by_deviceid(user_devid) + user = user_by_device_id(user_devid) if user: if "checkLogin" in request.path: _check_token( @@ -132,7 +132,7 @@ async def get_authcode(request: Request) -> Response: user_devid = request.query["deviceId"] # Ecovacs Home if user_devid: - user = user_by_deviceid(user_devid) + user = user_by_device_id(user_devid) if user: if "accessToken" in request.query: token = user_get_token(user["userid"], request.query["accessToken"]) @@ -224,8 +224,8 @@ def _auth_any( try: user_devid = devid countrycode = country - user = user_by_deviceid(user_devid) - bots = db_get().table("bots").all() + user = user_by_device_id(user_devid) + bots = _db_get().table("bots").all() login_details: EcoVacs_Login | EcoVacsHome_Login if user: # Default to user 0 diff --git a/bumper/web/middlewares.py b/bumper/web/middlewares.py index 0f4b127..760a9a5 100644 --- a/bumper/web/middlewares.py +++ b/bumper/web/middlewares.py @@ -14,7 +14,10 @@ _LOGGER = get_logger("webserver_requests") class CustomEncoder(json.JSONEncoder): + """Custom json encoder, which supports set.""" + def default(self, obj: Any) -> Any: + """Convert objects, which are not supported by the default JSONEncoder.""" if isinstance(obj, set): return list(obj) return json.JSONEncoder.default(self, obj) @@ -30,6 +33,7 @@ _EXCLUDE_FROM_LOGGING = [ @web.middleware async def log_all_requests(request: Request, handler: Handler) -> StreamResponse: + """Middleware to log all requests.""" if ( not request.match_info.route.resource ) or request.match_info.route.resource.canonical in _EXCLUDE_FROM_LOGGING: diff --git a/bumper/web/plugins/api/appsvr.py b/bumper/web/plugins/api/appsvr.py index 4fb0bf6..0b9016c 100644 --- a/bumper/web/plugins/api/appsvr.py +++ b/bumper/web/plugins/api/appsvr.py @@ -13,7 +13,7 @@ from aiohttp.web_routedef import AbstractRouteDef from amqtt.session import Session import bumper -from bumper.db import db_get, token_by_authcode, user_add_oauth +from bumper.db import _db_get, token_by_authcode, user_add_oauth from .. import WebserverPlugin from .pim import get_product_iot_map @@ -73,7 +73,7 @@ async def _handle_appsvr_app(request: Request) -> Response: todo = postbody["todo"] if todo == "GetGlobalDeviceList": - bots = db_get().table("bots").all() + bots = _db_get().table("bots").all() devices = [] for bot in bots: if bot["class"] != "": diff --git a/bumper/web/plugins/api/users.py b/bumper/web/plugins/api/users.py index 5cb2699..91d96b8 100644 --- a/bumper/web/plugins/api/users.py +++ b/bumper/web/plugins/api/users.py @@ -10,7 +10,13 @@ from aiohttp.web_response import Response from aiohttp.web_routedef import AbstractRouteDef from bumper import bumper_announce_ip -from bumper.db import bot_remove, bot_set_nick, check_authcode, db_get, loginByItToken +from bumper.db import ( + _db_get, + bot_remove, + bot_set_nick, + check_authcode, + login_by_it_token, +) from .. import WebserverPlugin @@ -81,7 +87,7 @@ async def _handle_usersapi(request: Request) -> Response: "userId": postbody["userId"], } else: # EcoVacs Home LoginByITToken - login_token = loginByItToken(postbody["token"]) + login_token = login_by_it_token(postbody["token"]) if login_token: body = { "resource": postbody["resource"], @@ -95,7 +101,7 @@ async def _handle_usersapi(request: Request) -> Response: elif todo == "GetDeviceList": body = { - "devices": db_get().table("bots").all(), + "devices": _db_get().table("bots").all(), "result": "ok", "todo": "result", } diff --git a/bumper/web/plugins/v1/private/user.py b/bumper/web/plugins/v1/private/user.py index aaeb7a7..c381262 100644 --- a/bumper/web/plugins/v1/private/user.py +++ b/bumper/web/plugins/v1/private/user.py @@ -8,7 +8,7 @@ from aiohttp.web_request import Request from aiohttp.web_response import Response from aiohttp.web_routedef import AbstractRouteDef -from bumper.db import check_token, user_by_deviceid, user_revoke_token +from bumper.db import check_token, user_by_device_id, user_revoke_token from bumper.web import auth_util from ... import WebserverPlugin, get_success_response @@ -84,7 +84,7 @@ async def _logout(request: Request) -> Response: try: user_device_id = request.match_info.get("devid", None) if user_device_id: - user = user_by_deviceid(user_device_id) + user = user_by_device_id(user_device_id) if user: if check_token(user["userid"], request.query["accessToken"]): # Deactivate old tokens and authcodes @@ -101,7 +101,7 @@ async def _logout(request: Request) -> Response: async def _get_user_account_info(request: Request) -> Response: try: user_devid = request.match_info.get("devid", "") - user = user_by_deviceid(user_devid) + user = user_by_device_id(user_devid) if user: username = f"fusername_{user['userid']}" return get_success_response( diff --git a/bumper/web/server.py b/bumper/web/server.py index f9da533..e207663 100644 --- a/bumper/web/server.py +++ b/bumper/web/server.py @@ -15,7 +15,7 @@ from aiohttp.web_request import Request from aiohttp.web_response import Response import bumper -from bumper.db import bot_get, bot_remove, client_get, client_remove, db_get +from bumper.db import _db_get, bot_get, bot_remove, client_get, client_remove from bumper.dns import get_resolver_with_public_nameserver from bumper.util import get_logger from bumper.web.middlewares import log_all_requests @@ -148,8 +148,8 @@ class WebServer: async def _handle_base(self, request: Request) -> Response: try: - bots = db_get().table("bots").all() - clients = db_get().table("clients").all() + bots = _db_get().table("bots").all() + clients = _db_get().table("clients").all() mq_sessions = [] for (session, _) in bumper.mqtt_server.broker._sessions.values(): mq_sessions.append( diff --git a/bumper/xmppserver.py b/bumper/xmppserver.py index f07c3d5..8461f92 100644 --- a/bumper/xmppserver.py +++ b/bumper/xmppserver.py @@ -1,3 +1,4 @@ +"""XMPP module.""" import asyncio import base64 import re @@ -23,6 +24,8 @@ boterrorlog = bumper.get_logger("boterror") class XMPPServer: + """XMPP server.""" + server_id = "ecouser.net" clients: list["XMPPAsyncClient"] = [] exit_flag = False @@ -35,6 +38,7 @@ class XMPPServer: self.xmpp_protocol = lambda: XMPPServer_Protocol() async def start_async_server(self) -> None: + """Start server.""" try: xmppserverlog.info(f"Starting XMPP Server at {self._host}:{self._port}") @@ -51,7 +55,7 @@ class XMPPServer: raise e def disconnect(self) -> None: - + """Disconnect.""" xmppserverlog.debug("waiting for all clients to disconnect") for client in self.clients: client._disconnect() @@ -62,11 +66,14 @@ class XMPPServer: class XMPPServer_Protocol(asyncio.Protocol): + """XMPP server protocol.""" + client_id = None exit_flag = False _client: Optional["XMPPAsyncClient"] = None def connection_made(self, transport: transports.BaseTransport) -> None: + """Establish connection.""" if self._client: # Existing client... upgrading to TLS xmppserverlog.debug(f"Upgraded connection for {self._client.address}") self._client.transport = transport @@ -78,6 +85,7 @@ class XMPPServer_Protocol(asyncio.Protocol): xmppserverlog.debug(f"New Connection from {client.address}") def connection_lost(self, exc: Exception | None) -> None: + """Lost connection.""" if self._client: XMPPServer.clients.remove(self._client) self._client.set_state("DISCONNECT") @@ -90,11 +98,14 @@ class XMPPServer_Protocol(asyncio.Protocol): ) def data_received(self, data: bytes) -> None: + """Parse received data.""" if self._client: self._client.parse_data(data) class XMPPAsyncClient: + """XMPP client.""" + IDLE = 0 CONNECT = 1 INIT = 2 @@ -120,6 +131,7 @@ class XMPPAsyncClient: xmppserverlog.debug(f"new client with ip {self.address}") def send(self, command: str) -> None: + """Send command.""" try: if self.log_sent_message: xmppserverlog.debug( @@ -158,6 +170,7 @@ class XMPPAsyncClient: return tag def set_state(self, state: str) -> None: + """Set state.""" try: new_state = getattr(XMPPAsyncClient, state) if self.state > new_state: @@ -240,7 +253,7 @@ class XMPPAsyncClient: and client.state == client.READY ): ctl_to = xml.get("to") - if not "from" in xml.attrib: + if "from" not in xml.attrib: xml.attrib["from"] = f"{self.bumper_jid}" rxmlstring = ET.tostring(xml).decode("utf-8") # clean up string to remove namespaces added by ET @@ -270,7 +283,7 @@ class XMPPAsyncClient: else: pingfrom = self.bumper_jid - if not "from" in xml.attrib: + if "from" not in xml.attrib: xml.attrib["from"] = f"{pingfrom}" pingstring = ET.tostring(xml).decode("utf-8") # clean up string to remove namespaces added by ET @@ -292,6 +305,7 @@ class XMPPAsyncClient: xmppserverlog.exception(f"{e}") async def schedule_ping(self, time: int) -> None: + """Schedule ping.""" if not self.state == 5: # disconnected pingstring = "".format( XMPPServer.server_id, self.bumper_jid @@ -303,7 +317,7 @@ class XMPPAsyncClient: def _handle_result(self, xml: ET.Element, data: str) -> None: try: ctl_to = xml.get("to") - if not "from" in xml.attrib: + if "from" not in xml.attrib: xml.attrib["from"] = f"{self.bumper_jid}" if "errno" in data: xmppserverlog.error(f"Error from bot - {data}") @@ -383,7 +397,7 @@ class XMPPAsyncClient: client.bumper_jid != self.bumper_jid and client.state == client.READY ): - if not "@" in ctl_to: # No user@, send to all clients? + if "@" not in ctl_to: # No user@, send to all clients? # TODO: Revisit later, this may be wrong client.send(rxmlstring) @@ -684,6 +698,7 @@ class XMPPAsyncClient: self.send(f' dummy ') def parse_data(self, data: bytes) -> None: + """Parse data.""" if data.decode("utf-8").startswith( " - client is signalling end of session/disconnect - if not "" in newdata: + if "" not in newdata: xmppserverlog.error(f"xml parse error - {newdata} - {e}") else: self.send("") # Close stream @@ -781,7 +796,7 @@ class XMPPAsyncClient: xmppserverlog.debug(f"Handling connect data - {newdata}") self._handle_connect(newdata.encode("utf-8")) else: - if not "" in newdata: + if "" not in newdata: xmppserverlog.error(f"xml parse error - {newdata} - {e}") else: self.send("") # Close stream diff --git a/tests/test_db.py b/tests/test_db.py index e8634cd..6a603e7 100644 --- a/tests/test_db.py +++ b/tests/test_db.py @@ -11,7 +11,7 @@ def test_db_path(): env = os.environ.copy() env.pop("DB_FILE") with mock.patch.dict(os.environ, env, clear=True): - assert db.db_file() == os.path.join(data_dir, "bumper.db") + assert db._db_file() == os.path.join(data_dir, "bumper.db") def test_user_db(): @@ -24,7 +24,7 @@ def test_user_db(): db.user_add_device("testuser", "dev_1234") # Add device to testuser assert ( - db.user_by_deviceid("dev_1234")["userid"] == "testuser" + db.user_by_device_id("dev_1234")["userid"] == "testuser" ) # Test that testuser was found by deviceid db.user_remove_device("testuser", "dev_1234") # Remove device from testuser