fix flake8 findings
This commit is contained in:
parent
23a84dda04
commit
a0eb86a76f
16 changed files with 230 additions and 121 deletions
|
|
@ -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
|
||||||
|
|
|
||||||
215
bumper/db.py
215
bumper/db.py
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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]
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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())
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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"] != "":
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue