diff --git a/Makefile b/Makefile index b6052c01f..e902b3592 100644 --- a/Makefile +++ b/Makefile @@ -50,6 +50,7 @@ test-wallets: test-unit: LNBITS_DATA_FOLDER="./tests/data" \ + LNBITS_DATABASE_URL="" \ LNBITS_BACKEND_WALLET_CLASS="FakeWallet" \ PYTHONUNBUFFERED=1 \ DEBUG=true \ @@ -57,6 +58,7 @@ test-unit: test-api: LNBITS_DATA_FOLDER="./tests/data" \ + LNBITS_DATABASE_URL="" \ LNBITS_BACKEND_WALLET_CLASS="FakeWallet" \ PYTHONUNBUFFERED=1 \ DEBUG=true \ diff --git a/lnbits/core/views/auth_api.py b/lnbits/core/views/auth_api.py index 94b6fa7da..234f8559c 100644 --- a/lnbits/core/views/auth_api.py +++ b/lnbits/core/views/auth_api.py @@ -67,7 +67,7 @@ async def login(data: LoginUsernamePassword) -> JSONResponse: raise HTTPException( status_code=HTTPStatus.UNAUTHORIZED, detail="Invalid credentials." ) - return _auth_success_response(account.username, account.id) + return _auth_success_response(account.username, account.id, account.email) @auth_router.post("/nostr", description="Login via Nostr") @@ -100,7 +100,7 @@ async def login_usr(data: LoginUsr) -> JSONResponse: raise HTTPException( status_code=HTTPStatus.UNAUTHORIZED, detail="User ID does not exist." ) - return _auth_success_response(account.username, account.id) + return _auth_success_response(account.username, account.id, account.email) @auth_router.get("/{provider}", description="SSO Provider") @@ -188,7 +188,7 @@ async def register(data: CreateUser) -> JSONResponse: ) account.hash_password(data.password) await create_account(account) - return _auth_success_response(account.username) + return _auth_success_response(account.username, account.id, account.email) @auth_router.put("/pubkey") @@ -271,22 +271,19 @@ async def reset_password(data: ResetUserPassword) -> JSONResponse: raise HTTPException( HTTPStatus.UNAUTHORIZED, "Auth by 'Username and Password' not allowed." ) - if not data.reset_key[:10].startswith("reset_key_"): - raise HTTPException(HTTPStatus.BAD_REQUEST, "This is not a reset key.") + + assert data.password == data.password_repeat, "Passwords do not match." + assert data.reset_key[:10].startswith("reset_key_"), "This is not a reset key." reset_data_json = decrypt_internal_message( base64.b64decode(data.reset_key[10:]).decode() ) - if not reset_data_json: - raise HTTPException(HTTPStatus.BAD_REQUEST, "Cannot process reset key.") + assert reset_data_json, "Cannot process reset key." action, user_id, request_time = json.loads(reset_data_json) - if not action: - raise HTTPException(HTTPStatus.BAD_REQUEST, "Missing action.") - if not user_id: - raise HTTPException(HTTPStatus.BAD_REQUEST, "Missing user ID.") - if not request_time: - raise HTTPException(HTTPStatus.BAD_REQUEST, "Missing reset time.") + assert action, "Missing action." + assert user_id, "Missing user ID." + assert request_time, "Missing reset time." _validate_auth_timeout(request_time) @@ -296,9 +293,7 @@ async def reset_password(data: ResetUserPassword) -> JSONResponse: account.hash_password(data.password) await update_account(account) - return _auth_success_response( - username=account.username, user_id=user_id, email=account.email - ) + return _auth_success_response(account.username, user_id, account.email) @auth_router.put("/update") @@ -365,7 +360,7 @@ async def first_install(data: UpdateSuperuserPassword) -> JSONResponse: account.hash_password(data.password) await update_account(account) settings.first_install = False - return _auth_success_response(username=account.username) + return _auth_success_response(account.username, account.id, account.email) async def _handle_sso_login(userinfo: OpenID, verified_user_id: Optional[str] = None): diff --git a/tests/api/test_auth.py b/tests/api/test_auth.py index 1604d9c2c..89e4596cc 100644 --- a/tests/api/test_auth.py +++ b/tests/api/test_auth.py @@ -95,6 +95,7 @@ async def test_login_alan_username_password_ok( payload: dict = jwt.decode(access_token, settings.auth_secret_key, ["HS256"]) access_token_payload = AccessTokenPayload(**payload) + assert access_token_payload.sub == "alan", "Subject is Alan." assert access_token_payload.email == "alan@lnbits.com" assert access_token_payload.auth_time, "Auth time should be set by server." @@ -113,7 +114,9 @@ async def test_login_alan_username_password_ok( assert not user.admin, "Not admin." assert not user.super_user, "Not superuser." assert user.has_password, "Password configured." - assert len(user.wallets) == 1, "One default wallet." + assert ( + len(user.wallets) == 1 + ), f"Expected 1 default wallet, not {len(user.wallets)}." @pytest.mark.asyncio @@ -221,7 +224,9 @@ async def test_register_ok(http_client: AsyncClient): assert not user.admin, "Not admin." assert not user.super_user, "Not superuser." assert user.has_password, "Password configured." - assert len(user.wallets) == 1, "One default wallet." + assert ( + len(user.wallets) == 1 + ), f"Expected 1 default wallet, not {len(user.wallets)}." @pytest.mark.asyncio @@ -509,7 +514,9 @@ async def test_register_nostr_ok(http_client: AsyncClient): assert not user.admin, "Not admin." assert not user.super_user, "Not superuser." assert not user.has_password, "Password configured." - assert len(user.wallets) == 1, "One default wallet." + assert ( + len(user.wallets) == 1 + ), f"Expected 1 default wallet, not {len(user.wallets)}." @pytest.mark.asyncio diff --git a/tests/conftest.py b/tests/conftest.py index 823c23a9e..4ccfaa3f9 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -102,7 +102,7 @@ async def db(): yield Database("database") -@pytest_asyncio.fixture(scope="package") +@pytest_asyncio.fixture(scope="session") async def user_alan(): account = await get_account_by_username("alan") if not account: @@ -112,6 +112,7 @@ async def user_alan(): username="alan", ) account.hash_password("secret1234") + account = await create_account(account) yield account