fix flake8 findings

This commit is contained in:
Robert Resch 2022-08-25 11:50:11 +02:00
parent 23a84dda04
commit a0eb86a76f
16 changed files with 230 additions and 121 deletions

View file

@ -1,3 +1,4 @@
"""Init module."""
import asyncio import asyncio
import logging import logging
import os import os
@ -5,8 +6,8 @@ import socket
import sys import sys
from bumper.db import ( from bumper.db import (
bot_reset_connectionStatus, bot_reset_connection_status,
client_reset_connectionStatus, client_reset_connection_status,
revoke_expired_oauths, revoke_expired_oauths,
revoke_expired_tokens, revoke_expired_tokens,
) )
@ -18,6 +19,7 @@ from bumper.xmppserver import XMPPServer
def strtobool(strbool: str | bool | None) -> bool: def strtobool(strbool: str | bool | None) -> bool:
"""Convert str to bool."""
if str(strbool).lower() in ["true", "1", "t", "y", "on", "yes"]: if str(strbool).lower() in ["true", "1", "t", "y", "on", "yes"]:
return True return True
else: else:
@ -77,13 +79,14 @@ web_server_bindings = [
async def start() -> None: async def start() -> None:
"""Start bumper."""
# Reset xmpp/mqtt to false in database for bots and clients # Reset xmpp/mqtt to false in database for bots and clients
bot_reset_connectionStatus() bot_reset_connection_status()
client_reset_connectionStatus() client_reset_connection_status()
try: try:
loop = asyncio.get_event_loop() loop = asyncio.get_event_loop()
except: except: # noqa: E722
loop = asyncio.new_event_loop() loop = asyncio.new_event_loop()
if bumper_debug: if bumper_debug:
@ -150,22 +153,28 @@ async def start() -> None:
async def maintenance() -> None: async def maintenance() -> None:
"""Run maintenance."""
revoke_expired_tokens() revoke_expired_tokens()
revoke_expired_oauths() revoke_expired_oauths()
async def shutdown() -> None: async def shutdown() -> None:
"""Shutdown bumper."""
try: try:
bumperlog.info("Shutting down") bumperlog.info("Shutting down")
global shutting_down global shutting_down
shutting_down = True shutting_down = True
global mqtt_helperbot
await mqtt_helperbot.disconnect() await mqtt_helperbot.disconnect()
global web_server
await web_server.shutdown() await web_server.shutdown()
global mqtt_server
while mqtt_server.state == "starting": while mqtt_server.state == "starting":
await asyncio.sleep(0.1) await asyncio.sleep(0.1)
if mqtt_server.state == "started": if mqtt_server.state == "started":
await mqtt_server.shutdown() await mqtt_server.shutdown()
global xmpp_server
if xmpp_server.server: if xmpp_server.server:
if xmpp_server.server.is_serving: if xmpp_server.server.is_serving:
xmpp_server.server.close() xmpp_server.server.close()
@ -177,6 +186,7 @@ async def shutdown() -> None:
def main(argv: None | list[str] = None) -> None: def main(argv: None | list[str] = None) -> None:
"""Start everything."""
import argparse import argparse
global bumper_debug global bumper_debug

View file

@ -1,3 +1,4 @@
"""Database module."""
import os import os
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import Any from typing import Any
@ -13,17 +14,17 @@ from .util import get_logger
_LOGGER = get_logger("db") _LOGGER = get_logger("db")
def db_file() -> str: def _db_file() -> str:
return os.environ.get("DB_FILE") or os_db_path() 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") 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 # 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 # Will create the tables if they don't exist
db.table("users", cache_size=0) db.table("users", cache_size=0)
@ -36,29 +37,32 @@ def db_get() -> TinyDB:
def user_add(userid: str) -> None: def user_add(userid: str) -> None:
"""Add user."""
newuser = BumperUser() newuser = BumperUser()
newuser.userid = userid newuser.userid = userid
user = user_get(userid) user = user_get(userid)
if not user: if not user:
_LOGGER.info(f"Adding new user with userid: {newuser.userid}") _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: def user_get(userid: str) -> None | Document:
users = db_get().table("users") """Get user."""
users = _db_get().table("users")
User = Query() User = Query()
return users.get(User.userid == userid) return users.get(User.userid == userid)
def user_by_deviceid(deviceid: str) -> None | Document: def user_by_device_id(deviceid: str) -> None | Document:
users = db_get().table("users") """Get user by device id."""
users = _db_get().table("users")
User = Query() User = Query()
return users.get(User.devices.any([deviceid])) return users.get(User.devices.any([deviceid]))
def user_full_upsert(user: dict[str, Any]) -> None: def _user_full_upsert(user: dict[str, Any]) -> None:
opendb = db_get() opendb = _db_get()
with opendb: with opendb:
users = opendb.table("users") users = opendb.table("users")
User = Query() User = Query()
@ -66,21 +70,23 @@ def user_full_upsert(user: dict[str, Any]) -> None:
def user_add_device(userid: str, devid: str) -> None: def user_add_device(userid: str, devid: str) -> None:
opendb = db_get() """Add device to user."""
opendb = _db_get()
with opendb: with opendb:
users = opendb.table("users") users = opendb.table("users")
User = Query() User = Query()
user = users.get(User.userid == userid) user = users.get(User.userid == userid)
if user: if user:
userdevices = list(user["devices"]) userdevices = list(user["devices"])
if not devid in userdevices: if devid not in userdevices:
userdevices.append(devid) userdevices.append(devid)
users.upsert({"devices": userdevices}, User.userid == userid) users.upsert({"devices": userdevices}, User.userid == userid)
def user_remove_device(userid: str, devid: str) -> None: def user_remove_device(userid: str, devid: str) -> None:
opendb = db_get() """Remove device from user."""
opendb = _db_get()
with opendb: with opendb:
users = opendb.table("users") users = opendb.table("users")
User = Query() User = Query()
@ -94,21 +100,23 @@ def user_remove_device(userid: str, devid: str) -> None:
def user_add_bot(userid: str, did: str) -> None: def user_add_bot(userid: str, did: str) -> None:
opendb = db_get() """Add bot to user."""
opendb = _db_get()
with opendb: with opendb:
users = opendb.table("users") users = opendb.table("users")
User = Query() User = Query()
user = users.get(User.userid == userid) user = users.get(User.userid == userid)
if user: if user:
userbots = list(user["bots"]) userbots = list(user["bots"])
if not did in userbots: if did not in userbots:
userbots.append(did) userbots.append(did)
users.upsert({"bots": userbots}, User.userid == userid) users.upsert({"bots": userbots}, User.userid == userid)
def user_remove_bot(userid: str, did: str) -> None: def user_remove_bot(userid: str, did: str) -> None:
opendb = db_get() """Remove bot from user."""
opendb = _db_get()
with opendb: with opendb:
users = opendb.table("users") users = opendb.table("users")
User = Query() User = Query()
@ -122,17 +130,20 @@ def user_remove_bot(userid: str, did: str) -> None:
def user_get_tokens(userid: str) -> list[Document]: 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) return tokens.search(Query().userid == userid)
def user_get_token(userid: str, token: str) -> Document | None: 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)) return tokens.get((Query().userid == userid) & (Query().token == token))
def user_add_token(userid: str, token: str) -> None: def user_add_token(userid: str, token: str) -> None:
opendb = db_get() """Ass token for given user."""
opendb = _db_get()
with opendb: with opendb:
tokens = opendb.table("tokens") tokens = opendb.table("tokens")
tmptoken = tokens.get((Query().userid == userid) & (Query().token == token)) 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: def user_revoke_all_tokens(userid: str) -> None:
opendb = db_get() """Revoke all tokens for given user."""
opendb = _db_get()
with opendb: with opendb:
tokens = opendb.table("tokens") tokens = opendb.table("tokens")
tsearch = tokens.search(Query().userid == userid) 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: def user_revoke_expired_tokens(userid: str) -> None:
opendb = db_get() """Revoke expired user tokens."""
opendb = _db_get()
with opendb: with opendb:
tokens = opendb.table("tokens") tokens = opendb.table("tokens")
tsearch = tokens.search(Query().userid == userid) 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: def user_revoke_token(userid: str, token: str) -> None:
opendb = db_get() """Revoke user token."""
opendb = _db_get()
with opendb: with opendb:
tokens = opendb.table("tokens") tokens = opendb.table("tokens")
tmptoken = tokens.get((Query().userid == userid) & (Query().token == token)) 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: def user_add_authcode(userid: str, token: str, authcode: str) -> None:
opendb = db_get() """Add user authcode."""
opendb = _db_get()
with opendb: with opendb:
tokens = opendb.table("tokens") tokens = opendb.table("tokens")
tmptoken = tokens.get((Query().userid == userid) & (Query().token == token)) 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: def user_revoke_authcode(userid: str, token: str) -> None:
opendb = db_get() """Revoke user authcode."""
opendb = _db_get()
with opendb: with opendb:
tokens = opendb.table("tokens") tokens = opendb.table("tokens")
tmptoken = tokens.get((Query().userid == userid) & (Query().token == token)) 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: def revoke_expired_oauths() -> None:
opendb = db_get() """Revoke expired oauths."""
opendb = _db_get()
with opendb: with opendb:
table = opendb.table("oauth") table = opendb.table("oauth")
entries = table.all() entries = table.all()
@ -217,7 +234,8 @@ def revoke_expired_oauths() -> None:
def user_revoke_expired_oauths(userid: str) -> None: def user_revoke_expired_oauths(userid: str) -> None:
opendb = db_get() """Revoke expired oauths by user."""
opendb = _db_get()
with opendb: with opendb:
table = opendb.table("oauth") table = opendb.table("oauth")
search = table.search(Query().userid == userid) 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: def user_add_oauth(userid: str) -> OAuth:
"""Add oauth for user."""
user_revoke_expired_oauths(userid) user_revoke_expired_oauths(userid)
opendb = db_get() opendb = _db_get()
with opendb: with opendb:
table = opendb.table("oauth") table = opendb.table("oauth")
entry = table.get(Query().userid == userid) 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: 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) return tokens.get(Query().authcode == authcode)
def get_disconnected_xmpp_clients() -> list[Document]: def get_disconnected_xmpp_clients() -> list[Document]:
clients = db_get().table("clients") """Get disconnected XMPP clients."""
Client = Query() clients = _db_get().table("clients")
return clients.search(Client.xmpp_connection == False) client = Query()
return clients.search(client.xmpp_connection == False) # noqa: E712
def check_authcode(uid: str, authcode: str) -> bool: def check_authcode(uid: str, authcode: str) -> bool:
"""Check authcode."""
_LOGGER.debug(f"Checking for authcode: {authcode}") _LOGGER.debug(f"Checking for authcode: {authcode}")
tokens = db_get().table("tokens") tokens = _db_get().table("tokens")
tmpauth = tokens.get( tmpauth = tokens.get(
(Query().authcode == authcode) (Query().authcode == authcode)
& ( # Match authcode & ( # Match authcode
@ -270,9 +292,10 @@ def check_authcode(uid: str, authcode: str) -> bool:
return False 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}") _LOGGER.debug(f"Checking for authcode: {authcode}")
tokens = db_get().table("tokens") tokens = _db_get().table("tokens")
tmpauth = tokens.get( tmpauth = tokens.get(
Query().authcode Query().authcode
== authcode == authcode
@ -288,8 +311,9 @@ def loginByItToken(authcode: str) -> dict[str, str]:
def check_token(uid: str, token: str) -> bool: def check_token(uid: str, token: str) -> bool:
"""Check token."""
_LOGGER.debug(f"Checking for token: {token}") _LOGGER.debug(f"Checking for token: {token}")
tokens = db_get().table("tokens") tokens = _db_get().table("tokens")
tmpauth = tokens.get( tmpauth = tokens.get(
(Query().token == token) (Query().token == token)
& ( # Match token & ( # Match token
@ -304,122 +328,137 @@ def check_token(uid: str, token: str) -> bool:
def revoke_expired_tokens() -> None: def revoke_expired_tokens() -> None:
tokens = db_get().table("tokens").all() """Revoke expired tokens."""
tokens = _db_get().table("tokens").all()
for i in tokens: for i in tokens:
if datetime.now() >= datetime.fromisoformat(i["expiration"]): if datetime.now() >= datetime.fromisoformat(i["expiration"]):
_LOGGER.debug("Removing token {} due to expiration".format(i["token"])) _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: def bot_add(sn: str, did: str, dev_class: str, resource: str, company: str) -> None:
newbot = VacBotDevice() """Add bot."""
newbot.did = did new_bot = VacBotDevice()
newbot.name = sn new_bot.did = did
newbot.vac_bot_device_class = devclass new_bot.name = sn
newbot.resource = resource new_bot.vac_bot_device_class = dev_class
newbot.company = company new_bot.resource = resource
new_bot.company = company
bot = bot_get(did) bot = bot_get(did)
if not bot: # Not existing bot in database if not bot: # Not existing bot in database
if ( 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 ): # try to prevent bad additions to the bot list
_LOGGER.info(f"Adding new bot with SN: {newbot.name} DID: {newbot.did}") _LOGGER.info(f"Adding new bot with SN: {new_bot.name} DID: {new_bot.did}")
bot_full_upsert(newbot.asdict()) bot_full_upsert(new_bot.asdict())
def bot_remove(did: str) -> None: def bot_remove(did: str) -> None:
bots = db_get().table("bots") """Remove bot."""
bots = _db_get().table("bots")
bot = bot_get(did) bot = bot_get(did)
if bot: if bot:
bots.remove(doc_ids=[bot.doc_id]) bots.remove(doc_ids=[bot.doc_id])
def bot_get(did: str) -> Document | None: def bot_get(did: str) -> Document | None:
bots = db_get().table("bots") """Get bot."""
Bot = Query() bots = _db_get().table("bots")
return bots.get(Bot.did == did) bot = Query()
return bots.get(bot.did == did)
def bot_full_upsert(vacbot: dict[str, Any]) -> None: def bot_full_upsert(vacbot: dict[str, Any]) -> None:
bots = db_get().table("bots") """Upsert bot."""
Bot = Query() bots = _db_get().table("bots")
bot = Query()
if "did" in vacbot: if "did" in vacbot:
bots.upsert(vacbot, Bot.did == vacbot["did"]) bots.upsert(vacbot, bot.did == vacbot["did"])
else: else:
_LOGGER.error(f"No DID in vacbot: {vacbot}") _LOGGER.error(f"No DID in vacbot: {vacbot}")
def bot_set_nick(did: str, nick: str) -> None: def bot_set_nick(did: str, nick: str) -> None:
bots = db_get().table("bots") """Bot set nickname."""
Bot = Query() bots = _db_get().table("bots")
bots.upsert({"nick": nick}, Bot.did == did) bot = Query()
bots.upsert({"nick": nick}, bot.did == did)
def bot_set_mqtt(did: str, mqtt: bool) -> None: def bot_set_mqtt(did: str, mqtt: bool) -> None:
bots = db_get().table("bots") """Bot ste MQTT status."""
Bot = Query() bots = _db_get().table("bots")
bots.upsert({"mqtt_connection": mqtt}, Bot.did == did) bot = Query()
bots.upsert({"mqtt_connection": mqtt}, bot.did == did)
def bot_set_xmpp(did: str, xmpp: bool) -> None: def bot_set_xmpp(did: str, xmpp: bool) -> None:
bots = db_get().table("bots") """Bot set XMPP status."""
Bot = Query() bots = _db_get().table("bots")
bots.upsert({"xmpp_connection": xmpp}, Bot.did == did) bot = Query()
bots.upsert({"xmpp_connection": xmpp}, bot.did == did)
def client_add(userid: str, realm: str, resource: str) -> None: def client_add(userid: str, realm: str, resource: str) -> None:
newclient = VacBotClient() """Add client."""
newclient.userid = userid new_client = VacBotClient()
newclient.realm = realm new_client.userid = userid
newclient.resource = resource new_client.realm = realm
new_client.resource = resource
client = client_get(resource) client = client_get(resource)
if not client: if not client:
_LOGGER.info(f"Adding new client with resource {newclient.resource}") _LOGGER.info(f"Adding new client with resource {new_client.resource}")
client_full_upsert(newclient.asdict()) _client_full_upsert(new_client.asdict())
def client_remove(resource: str) -> None: def client_remove(resource: str) -> None:
clients = db_get().table("clients") """Remove client."""
clients = _db_get().table("clients")
client = client_get(resource) client = client_get(resource)
if client: if client:
clients.remove(doc_ids=[client.doc_id]) clients.remove(doc_ids=[client.doc_id])
def client_get(resource: str) -> Document | None: def client_get(resource: str) -> Document | None:
clients = db_get().table("clients") """Get client by resource."""
Client = Query() clients = _db_get().table("clients")
return clients.get(Client.resource == resource) client = Query()
return clients.get(client.resource == resource)
def client_full_upsert(client: dict[str, Any]) -> None: def _client_full_upsert(client: dict[str, Any]) -> None:
clients = db_get().table("clients") clients = _db_get().table("clients")
Client = Query() client_query = Query()
clients.upsert(client, Client.resource == client["resource"]) clients.upsert(client, client_query.resource == client["resource"])
def client_set_mqtt(resource: str, mqtt: bool) -> None: def client_set_mqtt(resource: str, mqtt: bool) -> None:
clients = db_get().table("clients") """Client set MQTT status."""
Client = Query() clients = _db_get().table("clients")
clients.upsert({"mqtt_connection": mqtt}, Client.resource == resource) client = Query()
clients.upsert({"mqtt_connection": mqtt}, client.resource == resource)
def client_set_xmpp(resource: str, xmpp: bool) -> None: def client_set_xmpp(resource: str, xmpp: bool) -> None:
clients = db_get().table("clients") """Client set XMPP status."""
Client = Query() clients = _db_get().table("clients")
clients.upsert({"xmpp_connection": xmpp}, Client.resource == resource) client = Query()
clients.upsert({"xmpp_connection": xmpp}, client.resource == resource)
def bot_reset_connectionStatus() -> None: def bot_reset_connection_status() -> None:
bots = db_get().table("bots") """Reset all bot connection status."""
bots = _db_get().table("bots")
for bot in bots: for bot in bots:
bot_set_mqtt(bot["did"], False) bot_set_mqtt(bot["did"], False)
bot_set_xmpp(bot["did"], False) bot_set_xmpp(bot["did"], False)
def client_reset_connectionStatus() -> None: def client_reset_connection_status() -> None:
clients = db_get().table("clients") """Reset all client connection status."""
clients = _db_get().table("clients")
for client in clients: for client in clients:
client_set_mqtt(client["resource"], False) client_set_mqtt(client["resource"], False)
client_set_xmpp(client["resource"], False) client_set_xmpp(client["resource"], False)

View file

@ -1,11 +1,14 @@
"""Dns module."""
from aiohttp import AsyncResolver from aiohttp import AsyncResolver
def get_resolver_with_public_nameserver() -> AsyncResolver: def get_resolver_with_public_nameserver() -> AsyncResolver:
"""Get resolver."""
# requires aiodns # requires aiodns
return AsyncResolver(nameservers=["1.1.1.1", "8.8.8.8"]) return AsyncResolver(nameservers=["1.1.1.1", "8.8.8.8"])
async def resolve(host: str) -> str: async def resolve(host: str) -> str:
"""Resolve host."""
hosts = await get_resolver_with_public_nameserver().resolve(host) hosts = await get_resolver_with_public_nameserver().resolve(host)
return hosts[0]["host"] # type:ignore[no-any-return] return hosts[0]["host"] # type:ignore[no-any-return]

View file

@ -1,3 +1,4 @@
"""Models module."""
import json import json
import uuid import uuid
from datetime import datetime, timedelta from datetime import datetime, timedelta
@ -8,6 +9,8 @@ from bumper.util import convert_to_millis
class VacBotDevice: class VacBotDevice:
"""Vacuum device."""
def __init__( def __init__(
self, self,
did: str = "", did: str = "",
@ -27,6 +30,7 @@ class VacBotDevice:
self.xmpp_connection = False self.xmpp_connection = False
def asdict(self) -> dict[str, str | bool]: def asdict(self) -> dict[str, str | bool]:
"""Convert to dict."""
return { return {
"class": self.vac_bot_device_class, "class": self.vac_bot_device_class,
"company": self.company, "company": self.company,
@ -40,16 +44,21 @@ class VacBotDevice:
class BumperUser: class BumperUser:
"""Bumper user."""
def __init__(self, userid: str = ""): def __init__(self, userid: str = ""):
self.userid = userid self.userid = userid
self.devices: list[str] = [] self.devices: list[str] = []
self.bots: list[str] = [] self.bots: list[str] = []
def asdict(self) -> dict[str, Any]: def asdict(self) -> dict[str, Any]:
"""Convert to dict."""
return {"userid": self.userid, "devices": self.devices, "bots": self.bots} return {"userid": self.userid, "devices": self.devices, "bots": self.bots}
class GlobalVacBotDevice(VacBotDevice): # EcoVacs Home class GlobalVacBotDevice(VacBotDevice):
"""Global vacuum device."""
UILogicId = "" UILogicId = ""
ota = True ota = True
updateInfo = {"changeLog": "", "needUpdate": False} updateInfo = {"changeLog": "", "needUpdate": False}
@ -58,6 +67,8 @@ class GlobalVacBotDevice(VacBotDevice): # EcoVacs Home
class VacBotClient: class VacBotClient:
"""Vacuum client."""
def __init__(self, userid: str = "", realm: str = "", token: str = ""): def __init__(self, userid: str = "", realm: str = "", token: str = ""):
self.userid = userid self.userid = userid
self.realm = realm self.realm = realm
@ -66,6 +77,7 @@ class VacBotClient:
self.xmpp_connection = False self.xmpp_connection = False
def asdict(self) -> dict[str, Any]: def asdict(self) -> dict[str, Any]:
"""Convert to dict."""
return { return {
"userid": self.userid, "userid": self.userid,
"realm": self.realm, "realm": self.realm,
@ -76,6 +88,8 @@ class VacBotClient:
class EcoVacs_Login: class EcoVacs_Login:
"""Ecovacs login."""
accessToken = "" accessToken = ""
country = "" country = ""
email = "" email = ""
@ -83,16 +97,21 @@ class EcoVacs_Login:
username = "" username = ""
def toJSON(self) -> str: def toJSON(self) -> str:
"""Convert to json."""
return json.dumps(self, default=lambda o: o.__dict__, sort_keys=False) return json.dumps(self, default=lambda o: o.__dict__, sort_keys=False)
class EcoVacsHome_Login(EcoVacs_Login): class EcoVacsHome_Login(EcoVacs_Login):
"""Ecovacs home login."""
loginName = "" loginName = ""
mobile: str | None = "" mobile: str | None = ""
ucUid = "" ucUid = ""
class OAuth: class OAuth:
"""Oauth."""
access_token = "" access_token = ""
expire_at = "" expire_at = ""
refresh_token = "" refresh_token = ""
@ -103,6 +122,7 @@ class OAuth:
@classmethod @classmethod
def create_new(cls, userId: str) -> "OAuth": def create_new(cls, userId: str) -> "OAuth":
"""Create new."""
oauth = OAuth() oauth = OAuth()
oauth.userId = userId oauth.userId = userId
oauth.access_token = uuid.uuid4().hex oauth.access_token = uuid.uuid4().hex
@ -113,9 +133,11 @@ class OAuth:
return oauth return oauth
def toDB(self) -> dict: def toDB(self) -> dict:
"""Convert for db."""
return self.__dict__ return self.__dict__
def toResponse(self) -> dict: def toResponse(self) -> dict:
"""Convert to response."""
data = self.__dict__ data = self.__dict__
data["expire_at"] = convert_to_millis( data["expire_at"] = convert_to_millis(
datetime.fromisoformat(self.expire_at).timestamp() datetime.fromisoformat(self.expire_at).timestamp()

View file

@ -154,6 +154,7 @@ class HelperBot:
self._commands.pop(request_id, None) self._commands.pop(request_id, None)
def publish(self, topic: str, data: bytes) -> None: def publish(self, topic: str, data: bytes) -> None:
"""Publish message."""
self._client.publish(topic, data) self._client.publish(topic, data)
async def disconnect(self) -> None: async def disconnect(self) -> None:

View file

@ -1,6 +1,5 @@
"""Mqtt proxy module.""" """Mqtt proxy module."""
import asyncio import asyncio
import re
import ssl import ssl
import typing import typing
from collections.abc import MutableMapping from collections.abc import MutableMapping
@ -52,6 +51,7 @@ class ProxyClient:
self._port = port self._port = port
async def connect(self, username: str, password: str) -> None: async def connect(self, username: str, password: str) -> None:
"""Connect."""
try: try:
await self._client.connect( await self._client.connect(
f"mqtts://{username}:{password}@{self._host}:{self._port}" 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: 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)]) await self._client.subscribe([(topic, qos)])
async def disconnect(self) -> None: async def disconnect(self) -> None:
"""Disconnect."""
await self._client.disconnect() await self._client.disconnect()
async def publish(self, topic: str, message: bytes, qos: int | None = None) -> None: async def publish(self, topic: str, message: bytes, qos: int | None = None) -> None:
"""Publish message."""
await self._client.publish(topic, message, qos) await self._client.publish(topic, message, qos)

View file

@ -253,6 +253,8 @@ class BumperMQTTServerPlugin:
async def on_broker_client_subscribed( async def on_broker_client_subscribed(
self, client_id: str, topic: str, qos: QOS_0 | QOS_1 | QOS_2 self, client_id: str, topic: str, qos: QOS_0 | QOS_1 | QOS_2
) -> None: ) -> None:
"""Is called when a client subscribes on the broker."""
if bumper.bumper_proxy_mqtt: if bumper.bumper_proxy_mqtt:
# if proxy mode, also subscribe on ecovacs server # if proxy mode, also subscribe on ecovacs server
if client_id in self._proxy_clients: if client_id in self._proxy_clients:

View file

@ -1,3 +1,5 @@
"""Util module."""
import logging import logging
import os import os
import sys 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: def get_logger(name: str, rotate: RotatingFileHandler | None = None) -> logging.Logger:
"""Get logger."""
found_logger = __loggers.get(name) found_logger = __loggers.get(name)
if found_logger: if found_logger:
return found_logger return found_logger
@ -53,4 +56,5 @@ def convert_to_millis(seconds: int | float) -> int:
def get_current_time_as_millis() -> int: def get_current_time_as_millis() -> int:
"""Get current time in millis."""
return convert_to_millis(datetime.utcnow().timestamp()) return convert_to_millis(datetime.utcnow().timestamp())

View file

@ -11,13 +11,13 @@ from aiohttp.web_response import Response
from bumper import db, use_auth from bumper import db, use_auth
from bumper.db import ( from bumper.db import (
db_get, _db_get,
user_add, user_add,
user_add_authcode, user_add_authcode,
user_add_bot, user_add_bot,
user_add_device, user_add_device,
user_add_token, user_add_token,
user_by_deviceid, user_by_device_id,
user_get, user_get,
user_get_token, user_get_token,
user_revoke_expired_tokens, user_revoke_expired_tokens,
@ -60,7 +60,7 @@ async def login(request: Request) -> Response:
if ( if (
not user_devid == "" not user_devid == ""
): # Performing basic "auth" using devid, super insecure ): # Performing basic "auth" using devid, super insecure
user = user_by_deviceid(user_devid) user = user_by_device_id(user_devid)
if user: if user:
if "checkLogin" in request.path: if "checkLogin" in request.path:
_check_token( _check_token(
@ -132,7 +132,7 @@ async def get_authcode(request: Request) -> Response:
user_devid = request.query["deviceId"] # Ecovacs Home user_devid = request.query["deviceId"] # Ecovacs Home
if user_devid: if user_devid:
user = user_by_deviceid(user_devid) user = user_by_device_id(user_devid)
if user: if user:
if "accessToken" in request.query: if "accessToken" in request.query:
token = user_get_token(user["userid"], request.query["accessToken"]) token = user_get_token(user["userid"], request.query["accessToken"])
@ -224,8 +224,8 @@ def _auth_any(
try: try:
user_devid = devid user_devid = devid
countrycode = country countrycode = country
user = user_by_deviceid(user_devid) user = user_by_device_id(user_devid)
bots = db_get().table("bots").all() bots = _db_get().table("bots").all()
login_details: EcoVacs_Login | EcoVacsHome_Login login_details: EcoVacs_Login | EcoVacsHome_Login
if user: # Default to user 0 if user: # Default to user 0

View file

@ -14,7 +14,10 @@ _LOGGER = get_logger("webserver_requests")
class CustomEncoder(json.JSONEncoder): class CustomEncoder(json.JSONEncoder):
"""Custom json encoder, which supports set."""
def default(self, obj: Any) -> Any: def default(self, obj: Any) -> Any:
"""Convert objects, which are not supported by the default JSONEncoder."""
if isinstance(obj, set): if isinstance(obj, set):
return list(obj) return list(obj)
return json.JSONEncoder.default(self, obj) return json.JSONEncoder.default(self, obj)
@ -30,6 +33,7 @@ _EXCLUDE_FROM_LOGGING = [
@web.middleware @web.middleware
async def log_all_requests(request: Request, handler: Handler) -> StreamResponse: async def log_all_requests(request: Request, handler: Handler) -> StreamResponse:
"""Middleware to log all requests."""
if ( if (
not request.match_info.route.resource not request.match_info.route.resource
) or request.match_info.route.resource.canonical in _EXCLUDE_FROM_LOGGING: ) or request.match_info.route.resource.canonical in _EXCLUDE_FROM_LOGGING:

View file

@ -13,7 +13,7 @@ from aiohttp.web_routedef import AbstractRouteDef
from amqtt.session import Session from amqtt.session import Session
import bumper 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 .. import WebserverPlugin
from .pim import get_product_iot_map from .pim import get_product_iot_map
@ -73,7 +73,7 @@ async def _handle_appsvr_app(request: Request) -> Response:
todo = postbody["todo"] todo = postbody["todo"]
if todo == "GetGlobalDeviceList": if todo == "GetGlobalDeviceList":
bots = db_get().table("bots").all() bots = _db_get().table("bots").all()
devices = [] devices = []
for bot in bots: for bot in bots:
if bot["class"] != "": if bot["class"] != "":

View file

@ -10,7 +10,13 @@ from aiohttp.web_response import Response
from aiohttp.web_routedef import AbstractRouteDef from aiohttp.web_routedef import AbstractRouteDef
from bumper import bumper_announce_ip 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 from .. import WebserverPlugin
@ -81,7 +87,7 @@ async def _handle_usersapi(request: Request) -> Response:
"userId": postbody["userId"], "userId": postbody["userId"],
} }
else: # EcoVacs Home LoginByITToken else: # EcoVacs Home LoginByITToken
login_token = loginByItToken(postbody["token"]) login_token = login_by_it_token(postbody["token"])
if login_token: if login_token:
body = { body = {
"resource": postbody["resource"], "resource": postbody["resource"],
@ -95,7 +101,7 @@ async def _handle_usersapi(request: Request) -> Response:
elif todo == "GetDeviceList": elif todo == "GetDeviceList":
body = { body = {
"devices": db_get().table("bots").all(), "devices": _db_get().table("bots").all(),
"result": "ok", "result": "ok",
"todo": "result", "todo": "result",
} }

View file

@ -8,7 +8,7 @@ from aiohttp.web_request import Request
from aiohttp.web_response import Response from aiohttp.web_response import Response
from aiohttp.web_routedef import AbstractRouteDef 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 bumper.web import auth_util
from ... import WebserverPlugin, get_success_response from ... import WebserverPlugin, get_success_response
@ -84,7 +84,7 @@ async def _logout(request: Request) -> Response:
try: try:
user_device_id = request.match_info.get("devid", None) user_device_id = request.match_info.get("devid", None)
if user_device_id: if user_device_id:
user = user_by_deviceid(user_device_id) user = user_by_device_id(user_device_id)
if user: if user:
if check_token(user["userid"], request.query["accessToken"]): if check_token(user["userid"], request.query["accessToken"]):
# Deactivate old tokens and authcodes # Deactivate old tokens and authcodes
@ -101,7 +101,7 @@ async def _logout(request: Request) -> Response:
async def _get_user_account_info(request: Request) -> Response: async def _get_user_account_info(request: Request) -> Response:
try: try:
user_devid = request.match_info.get("devid", "") user_devid = request.match_info.get("devid", "")
user = user_by_deviceid(user_devid) user = user_by_device_id(user_devid)
if user: if user:
username = f"fusername_{user['userid']}" username = f"fusername_{user['userid']}"
return get_success_response( return get_success_response(

View file

@ -15,7 +15,7 @@ from aiohttp.web_request import Request
from aiohttp.web_response import Response from aiohttp.web_response import Response
import bumper 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.dns import get_resolver_with_public_nameserver
from bumper.util import get_logger from bumper.util import get_logger
from bumper.web.middlewares import log_all_requests from bumper.web.middlewares import log_all_requests
@ -148,8 +148,8 @@ class WebServer:
async def _handle_base(self, request: Request) -> Response: async def _handle_base(self, request: Request) -> Response:
try: try:
bots = db_get().table("bots").all() bots = _db_get().table("bots").all()
clients = db_get().table("clients").all() clients = _db_get().table("clients").all()
mq_sessions = [] mq_sessions = []
for (session, _) in bumper.mqtt_server.broker._sessions.values(): for (session, _) in bumper.mqtt_server.broker._sessions.values():
mq_sessions.append( mq_sessions.append(

View file

@ -1,3 +1,4 @@
"""XMPP module."""
import asyncio import asyncio
import base64 import base64
import re import re
@ -23,6 +24,8 @@ boterrorlog = bumper.get_logger("boterror")
class XMPPServer: class XMPPServer:
"""XMPP server."""
server_id = "ecouser.net" server_id = "ecouser.net"
clients: list["XMPPAsyncClient"] = [] clients: list["XMPPAsyncClient"] = []
exit_flag = False exit_flag = False
@ -35,6 +38,7 @@ class XMPPServer:
self.xmpp_protocol = lambda: XMPPServer_Protocol() self.xmpp_protocol = lambda: XMPPServer_Protocol()
async def start_async_server(self) -> None: async def start_async_server(self) -> None:
"""Start server."""
try: try:
xmppserverlog.info(f"Starting XMPP Server at {self._host}:{self._port}") xmppserverlog.info(f"Starting XMPP Server at {self._host}:{self._port}")
@ -51,7 +55,7 @@ class XMPPServer:
raise e raise e
def disconnect(self) -> None: def disconnect(self) -> None:
"""Disconnect."""
xmppserverlog.debug("waiting for all clients to disconnect") xmppserverlog.debug("waiting for all clients to disconnect")
for client in self.clients: for client in self.clients:
client._disconnect() client._disconnect()
@ -62,11 +66,14 @@ class XMPPServer:
class XMPPServer_Protocol(asyncio.Protocol): class XMPPServer_Protocol(asyncio.Protocol):
"""XMPP server protocol."""
client_id = None client_id = None
exit_flag = False exit_flag = False
_client: Optional["XMPPAsyncClient"] = None _client: Optional["XMPPAsyncClient"] = None
def connection_made(self, transport: transports.BaseTransport) -> None: def connection_made(self, transport: transports.BaseTransport) -> None:
"""Establish connection."""
if self._client: # Existing client... upgrading to TLS if self._client: # Existing client... upgrading to TLS
xmppserverlog.debug(f"Upgraded connection for {self._client.address}") xmppserverlog.debug(f"Upgraded connection for {self._client.address}")
self._client.transport = transport self._client.transport = transport
@ -78,6 +85,7 @@ class XMPPServer_Protocol(asyncio.Protocol):
xmppserverlog.debug(f"New Connection from {client.address}") xmppserverlog.debug(f"New Connection from {client.address}")
def connection_lost(self, exc: Exception | None) -> None: def connection_lost(self, exc: Exception | None) -> None:
"""Lost connection."""
if self._client: if self._client:
XMPPServer.clients.remove(self._client) XMPPServer.clients.remove(self._client)
self._client.set_state("DISCONNECT") self._client.set_state("DISCONNECT")
@ -90,11 +98,14 @@ class XMPPServer_Protocol(asyncio.Protocol):
) )
def data_received(self, data: bytes) -> None: def data_received(self, data: bytes) -> None:
"""Parse received data."""
if self._client: if self._client:
self._client.parse_data(data) self._client.parse_data(data)
class XMPPAsyncClient: class XMPPAsyncClient:
"""XMPP client."""
IDLE = 0 IDLE = 0
CONNECT = 1 CONNECT = 1
INIT = 2 INIT = 2
@ -120,6 +131,7 @@ class XMPPAsyncClient:
xmppserverlog.debug(f"new client with ip {self.address}") xmppserverlog.debug(f"new client with ip {self.address}")
def send(self, command: str) -> None: def send(self, command: str) -> None:
"""Send command."""
try: try:
if self.log_sent_message: if self.log_sent_message:
xmppserverlog.debug( xmppserverlog.debug(
@ -158,6 +170,7 @@ class XMPPAsyncClient:
return tag return tag
def set_state(self, state: str) -> None: def set_state(self, state: str) -> None:
"""Set state."""
try: try:
new_state = getattr(XMPPAsyncClient, state) new_state = getattr(XMPPAsyncClient, state)
if self.state > new_state: if self.state > new_state:
@ -240,7 +253,7 @@ class XMPPAsyncClient:
and client.state == client.READY and client.state == client.READY
): ):
ctl_to = xml.get("to") ctl_to = xml.get("to")
if not "from" in xml.attrib: if "from" not in xml.attrib:
xml.attrib["from"] = f"{self.bumper_jid}" xml.attrib["from"] = f"{self.bumper_jid}"
rxmlstring = ET.tostring(xml).decode("utf-8") rxmlstring = ET.tostring(xml).decode("utf-8")
# clean up string to remove namespaces added by ET # clean up string to remove namespaces added by ET
@ -270,7 +283,7 @@ class XMPPAsyncClient:
else: else:
pingfrom = self.bumper_jid pingfrom = self.bumper_jid
if not "from" in xml.attrib: if "from" not in xml.attrib:
xml.attrib["from"] = f"{pingfrom}" xml.attrib["from"] = f"{pingfrom}"
pingstring = ET.tostring(xml).decode("utf-8") pingstring = ET.tostring(xml).decode("utf-8")
# clean up string to remove namespaces added by ET # clean up string to remove namespaces added by ET
@ -292,6 +305,7 @@ class XMPPAsyncClient:
xmppserverlog.exception(f"{e}") xmppserverlog.exception(f"{e}")
async def schedule_ping(self, time: int) -> None: async def schedule_ping(self, time: int) -> None:
"""Schedule ping."""
if not self.state == 5: # disconnected if not self.state == 5: # disconnected
pingstring = "<iq from='{}' to='{}' id='s2c1' type='get'><ping xmlns='urn:xmpp:ping'/></iq>".format( pingstring = "<iq from='{}' to='{}' id='s2c1' type='get'><ping xmlns='urn:xmpp:ping'/></iq>".format(
XMPPServer.server_id, self.bumper_jid XMPPServer.server_id, self.bumper_jid
@ -303,7 +317,7 @@ class XMPPAsyncClient:
def _handle_result(self, xml: ET.Element, data: str) -> None: def _handle_result(self, xml: ET.Element, data: str) -> None:
try: try:
ctl_to = xml.get("to") ctl_to = xml.get("to")
if not "from" in xml.attrib: if "from" not in xml.attrib:
xml.attrib["from"] = f"{self.bumper_jid}" xml.attrib["from"] = f"{self.bumper_jid}"
if "errno" in data: if "errno" in data:
xmppserverlog.error(f"Error from bot - {data}") xmppserverlog.error(f"Error from bot - {data}")
@ -383,7 +397,7 @@ class XMPPAsyncClient:
client.bumper_jid != self.bumper_jid client.bumper_jid != self.bumper_jid
and client.state == client.READY 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 # TODO: Revisit later, this may be wrong
client.send(rxmlstring) client.send(rxmlstring)
@ -684,6 +698,7 @@ class XMPPAsyncClient:
self.send(f'<presence to="{self.bumper_jid}"> dummy </presence>') self.send(f'<presence to="{self.bumper_jid}"> dummy </presence>')
def parse_data(self, data: bytes) -> None: def parse_data(self, data: bytes) -> None:
"""Parse data."""
if data.decode("utf-8").startswith( if data.decode("utf-8").startswith(
"<?xml" "<?xml"
@ -770,7 +785,7 @@ class XMPPAsyncClient:
elif "not well-formed (invalid token)" in e.msg: elif "not well-formed (invalid token)" in e.msg:
# If a lone </stream:stream> - client is signalling end of session/disconnect # If a lone </stream:stream> - client is signalling end of session/disconnect
if not "</stream:stream>" in newdata: if "</stream:stream>" not in newdata:
xmppserverlog.error(f"xml parse error - {newdata} - {e}") xmppserverlog.error(f"xml parse error - {newdata} - {e}")
else: else:
self.send("</stream:stream>") # Close stream self.send("</stream:stream>") # Close stream
@ -781,7 +796,7 @@ class XMPPAsyncClient:
xmppserverlog.debug(f"Handling connect data - {newdata}") xmppserverlog.debug(f"Handling connect data - {newdata}")
self._handle_connect(newdata.encode("utf-8")) self._handle_connect(newdata.encode("utf-8"))
else: else:
if not "</stream:stream>" in newdata: if "</stream:stream>" not in newdata:
xmppserverlog.error(f"xml parse error - {newdata} - {e}") xmppserverlog.error(f"xml parse error - {newdata} - {e}")
else: else:
self.send("</stream:stream>") # Close stream self.send("</stream:stream>") # Close stream

View file

@ -11,7 +11,7 @@ def test_db_path():
env = os.environ.copy() env = os.environ.copy()
env.pop("DB_FILE") env.pop("DB_FILE")
with mock.patch.dict(os.environ, env, clear=True): 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(): def test_user_db():
@ -24,7 +24,7 @@ def test_user_db():
db.user_add_device("testuser", "dev_1234") # Add device to testuser db.user_add_device("testuser", "dev_1234") # Add device to testuser
assert ( 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 ) # Test that testuser was found by deviceid
db.user_remove_device("testuser", "dev_1234") # Remove device from testuser db.user_remove_device("testuser", "dev_1234") # Remove device from testuser