Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 25 additions & 29 deletions src/ssh_manager.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
import asyncio
import asyncssh
import inspect
import logging
import async_timeout
from typing import Any
from asyncssh import PermissionDenied
from tenacity import retry, stop_after_attempt, wait_fixed, retry_if_exception
from .database import get_all_servers
Expand Down Expand Up @@ -114,11 +116,7 @@ async def run_command(self, alias: str, command: str, timeout: float = COMMAND_T
# Re-raise the exception to be handled by the global error handler
raise
finally:
# --- Definitive Crash Fix ---
# A simple `if conn:` check is the safest way to prevent an `await`
# on a `None` object, regardless of how the connection failed.
if conn:
await conn.close()
await self._close_conn(conn)

async def kill_process(self, alias: str, pid: int) -> None:
"""Kills a process on a remote server."""
Expand All @@ -129,14 +127,7 @@ async def kill_process(self, alias: str, pid: int) -> None:
except Exception:
raise
finally:
logger.debug(f"In finally block for kill_process. Connection object is: {conn}")
if conn and not conn.is_closed():
logger.debug("Connection is valid and not closed, closing now.")
await conn.close()
elif conn:
logger.debug("Connection is already closed or closing.")
else:
logger.debug("Connection is None, nothing to close.")
await self._close_conn(conn)

async def start_shell_session(self, alias: str) -> None:
"""
Expand Down Expand Up @@ -198,14 +189,7 @@ async def download_file(self, alias: str, remote_path: str, local_path: str) ->
except Exception:
raise
finally:
logger.debug(f"In finally block for download_file. Connection object is: {conn}")
if conn and not conn.is_closed():
logger.debug("Connection is valid and not closed, closing now.")
await conn.close()
elif conn:
logger.debug("Connection is already closed or closing.")
else:
logger.debug("Connection is None, nothing to close.")
await self._close_conn(conn)

async def upload_file(self, alias: str, local_path: str, remote_path: str) -> None:
"""Uploads a file to a remote server."""
Expand All @@ -217,15 +201,27 @@ async def upload_file(self, alias: str, local_path: str, remote_path: str) -> No
except Exception:
raise
finally:
logger.debug(f"In finally block for upload_file. Connection object is: {conn}")
if conn and not conn.is_closed():
logger.debug("Connection is valid and not closed, closing now.")
await conn.close()
elif conn:
logger.debug("Connection is already closed or closing.")
else:
logger.debug("Connection is None, nothing to close.")
await self._close_conn(conn)

# --- Health Check (No longer needed) ---
# The start_health_check and stop_health_check methods are removed as they
# are not required with the new just-in-time connection model.

async def _close_conn(self, conn: Any) -> None:
"""
Safely closes an SSH connection, handling various library patterns.
"""
if conn is None:
return

close = getattr(conn, "close", None)
close_result = None
if callable(close):
close_result = close()

if inspect.isawaitable(close_result):
await close_result

wait_closed = getattr(conn, "wait_closed", None)
if callable(wait_closed):
await wait_closed()
105 changes: 105 additions & 0 deletions tests/test_bot_e2e.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,105 @@
import pytest
from unittest.mock import AsyncMock, MagicMock

# --- Mocks for Telegram Bot and SSH Manager ---

# Mock the SSHManager to control its behavior in tests
class MockSSHManager:
def __init__(self, conn):
self._conn = conn
self._close_conn_mock = AsyncMock()

async def _close_conn(self, conn):
await self._close_conn_mock(conn)
# In a real scenario, this would call the actual close logic
if conn:
if hasattr(conn, 'wait_closed'):
conn.close()
await conn.wait_closed()
else:
conn.close()


async def run_command(self, alias, command):
try:
# Simulate a successful connection and command execution
yield ("output line 1", "stdout")
finally:
await self._close_conn(self._conn)

# Mock connection classes from the unit tests
class SyncCloseConn:
def __init__(self):
self.closed = False
def close(self):
self.closed = True

class AsyncSSHLikeConn:
def __init__(self):
self.closed = False
self.waited = False
def close(self):
self.closed = True
async def wait_closed(self):
self.waited = True

# --- E2E Test ---

@pytest.mark.asyncio
@pytest.mark.parametrize("conn_type", [SyncCloseConn, AsyncSSHLikeConn])
async def test_handle_server_connection_e2e(mocker, conn_type):
"""
E2E test to ensure `handle_server_connection` completes without TypeError.
"""
# 1. Setup Mocks

# Import the function to test and the config object
from src.main import handle_server_connection
from src import main

# Mock the global ssh_manager instance used by the handler
mock_conn = conn_type()
mock_ssh_manager = MockSSHManager(conn=mock_conn)
mocker.patch('src.main.ssh_manager', mock_ssh_manager)

# Mock Telegram's update and context objects
update = AsyncMock()
context = AsyncMock()

# Configure the mock objects with necessary attributes
user_id = 12345
update.effective_user.id = user_id
update.effective_chat.id = user_id
update.callback_query.data = "connect_server_some_alias"
update.callback_query.answer = AsyncMock()
update.callback_query.message = AsyncMock()

# Patch the config to authorize the user
mocker.patch.object(main.config, 'whitelisted_users', [user_id])


# 2. Execute the handler
try:
await handle_server_connection(update, context)
except TypeError as e:
# We want to fail on the specific "await None" error, but not others
if "object NoneType can't be used in 'await' expression" in str(e):
pytest.fail("The TypeError related to 'await None' should not have occurred.")
except Exception as e:
pytest.fail(f"An unexpected exception occurred: {e}")

# 3. Assertions
# Verify that the connection was closed correctly
mock_ssh_manager._close_conn_mock.assert_awaited_once_with(mock_conn)

if isinstance(mock_conn, SyncCloseConn):
assert mock_conn.closed
elif isinstance(mock_conn, AsyncSSHLikeConn):
assert mock_conn.closed
assert mock_conn.waited

# Verify bot interactions
update.callback_query.edit_message_text.assert_called()
# The text is in the second call to edit_message_text
final_call_args = update.callback_query.edit_message_text.call_args_list[1]
assert "Connected to" in final_call_args.args[0]
Loading
Loading