Add terminal session management with heartbeat monitoring, idle timeout detection, session reuse logic, and command history panel UI with search and filtering capabilities
Tests / Backend Tests (Python) (3.10) (push) Has been cancelled
Tests / Backend Tests (Python) (3.11) (push) Has been cancelled
Tests / Backend Tests (Python) (3.12) (push) Has been cancelled
Tests / Frontend Tests (JS) (push) Has been cancelled
Tests / Integration Tests (push) Has been cancelled
Tests / All Tests Passed (push) Has been cancelled

This commit is contained in:
2025-12-18 13:49:40 -05:00
parent 493668f746
commit 5bc12d0729
74 changed files with 14849 additions and 78 deletions
+284
View File
@@ -0,0 +1,284 @@
"""
Unit tests for the Command Policy Engine.
Tests cover:
- Blocklist patterns (sensitive commands)
- Allowlist patterns (safe commands)
- Masking of sensitive values
- Policy evaluation logic
"""
import pytest
from app.security.command_policy import (
CommandPolicy,
CommandPolicyResult,
PolicyDecision,
evaluate_command,
)
class TestCommandPolicyBlocklist:
"""Test blocklist patterns that should be blocked."""
@pytest.fixture
def policy(self):
return CommandPolicy()
def test_blocks_password_keyword(self, policy):
result = policy.evaluate("echo my password is secret")
assert result.is_blocked
assert result.decision == PolicyDecision.BLOCK
def test_blocks_token_keyword(self, policy):
result = policy.evaluate("export TOKEN=abc123")
assert result.is_blocked
def test_blocks_docker_login(self, policy):
result = policy.evaluate("docker login -u user -p secret")
assert result.is_blocked
def test_blocks_curl_with_auth_header(self, policy):
result = policy.evaluate('curl -H "Authorization: Bearer xyz" https://api.com')
assert result.is_blocked
def test_blocks_wget_with_auth(self, policy):
result = policy.evaluate("wget --header='Authorization: token123' https://api.com")
assert result.is_blocked
def test_blocks_cat_ssh_key(self, policy):
result = policy.evaluate("cat ~/.ssh/id_rsa")
assert result.is_blocked
def test_blocks_cat_shadow_file(self, policy):
result = policy.evaluate("cat /etc/shadow")
assert result.is_blocked
def test_blocks_export_secret_var(self, policy):
result = policy.evaluate("export AWS_SECRET_KEY=abc123")
assert result.is_blocked
def test_blocks_sshpass(self, policy):
result = policy.evaluate("sshpass -p 'password' ssh user@host")
assert result.is_blocked
def test_blocks_mysql_with_password(self, policy):
result = policy.evaluate("mysql -u root -pMyPassword db")
assert result.is_blocked
def test_blocks_kubectl_get_secret(self, policy):
result = policy.evaluate("kubectl get secret my-secret -o yaml")
assert result.is_blocked
def test_blocks_ansible_vault(self, policy):
result = policy.evaluate("ansible-vault decrypt secrets.yml")
assert result.is_blocked
class TestCommandPolicyAllowlist:
"""Test allowlist patterns that should be allowed."""
@pytest.fixture
def policy(self):
return CommandPolicy()
def test_allows_ls(self, policy):
result = policy.evaluate("ls -la /var/log")
assert result.should_log
assert result.decision == PolicyDecision.ALLOW
def test_allows_cd(self, policy):
result = policy.evaluate("cd /home/user")
assert result.should_log
def test_allows_pwd(self, policy):
result = policy.evaluate("pwd")
assert result.should_log
def test_allows_whoami(self, policy):
result = policy.evaluate("whoami")
assert result.should_log
def test_allows_df(self, policy):
result = policy.evaluate("df -h")
assert result.should_log
def test_allows_free(self, policy):
result = policy.evaluate("free -m")
assert result.should_log
def test_allows_systemctl_status(self, policy):
result = policy.evaluate("systemctl status nginx")
assert result.should_log
def test_allows_systemctl_restart(self, policy):
result = policy.evaluate("systemctl restart docker")
assert result.should_log
def test_allows_journalctl(self, policy):
result = policy.evaluate("journalctl -u nginx -n 100")
assert result.should_log
def test_allows_docker_ps(self, policy):
result = policy.evaluate("docker ps -a")
assert result.should_log
def test_allows_docker_logs(self, policy):
result = policy.evaluate("docker logs container_name")
assert result.should_log
def test_allows_docker_compose_ps(self, policy):
result = policy.evaluate("docker compose ps")
assert result.should_log
def test_allows_tail(self, policy):
result = policy.evaluate("tail -f /var/log/syslog")
assert result.should_log
def test_allows_grep(self, policy):
result = policy.evaluate("grep error /var/log/nginx/error.log")
assert result.should_log
def test_allows_ip_addr(self, policy):
result = policy.evaluate("ip addr show")
assert result.should_log
def test_allows_ping(self, policy):
result = policy.evaluate("ping -c 4 google.com")
assert result.should_log
def test_allows_apt_list(self, policy):
result = policy.evaluate("apt list --installed")
assert result.should_log
def test_allows_git_status(self, policy):
result = policy.evaluate("git status")
assert result.should_log
def test_allows_zfs_list(self, policy):
result = policy.evaluate("zfs list")
assert result.should_log
def test_allows_clear(self, policy):
result = policy.evaluate("clear")
assert result.should_log
def test_allows_exit(self, policy):
result = policy.evaluate("exit")
assert result.should_log
class TestCommandPolicyMasking:
"""Test masking of sensitive values in allowed commands."""
@pytest.fixture
def policy(self):
return CommandPolicy()
def test_masks_password_flag(self, policy):
# This would be in allowlist if not for the password
result = policy.evaluate("some-tool --password=secret123")
# Should be blocked due to password keyword
assert result.is_blocked
def test_masks_url_credentials(self, policy):
# Test that URL credentials would be masked if command was allowed
policy_permissive = CommandPolicy(mode="permissive")
result = policy_permissive.evaluate("git clone https://user:[email protected]/repo.git")
if result.should_log:
assert "***" in result.masked_command
def test_preserves_safe_command(self, policy):
result = policy.evaluate("ls -la")
assert result.should_log
assert result.masked_command == "ls -la"
class TestCommandPolicyUnknown:
"""Test commands not in allowlist."""
@pytest.fixture
def policy(self):
return CommandPolicy()
def test_unknown_command_strict_mode(self, policy):
result = policy.evaluate("some-random-command --flag")
assert result.decision == PolicyDecision.UNKNOWN
assert not result.should_log
def test_unknown_command_permissive_mode(self):
policy = CommandPolicy(mode="permissive")
result = policy.evaluate("some-random-command --flag")
assert result.decision == PolicyDecision.ALLOW
assert result.should_log
class TestCommandPolicyHash:
"""Test command hashing for deduplication."""
@pytest.fixture
def policy(self):
return CommandPolicy()
def test_same_command_same_hash(self, policy):
result1 = policy.evaluate("ls -la")
result2 = policy.evaluate("ls -la")
assert result1.command_hash == result2.command_hash
def test_normalized_whitespace_same_hash(self, policy):
result1 = policy.evaluate("ls -la")
result2 = policy.evaluate("ls -la")
assert result1.command_hash == result2.command_hash
def test_different_commands_different_hash(self, policy):
result1 = policy.evaluate("ls -la")
result2 = policy.evaluate("ls -l")
assert result1.command_hash != result2.command_hash
class TestCommandPolicyEdgeCases:
"""Test edge cases and special inputs."""
@pytest.fixture
def policy(self):
return CommandPolicy()
def test_empty_command(self, policy):
result = policy.evaluate("")
assert result.decision == PolicyDecision.UNKNOWN
assert not result.should_log
def test_whitespace_only(self, policy):
result = policy.evaluate(" ")
assert result.decision == PolicyDecision.UNKNOWN
def test_case_insensitive_blocklist(self, policy):
result = policy.evaluate("export PASSWORD=secret")
assert result.is_blocked
result = policy.evaluate("export password=secret")
assert result.is_blocked
def test_convenience_functions(self):
result = evaluate_command("ls -la")
assert result.should_log
class TestCommandPolicyResult:
"""Test CommandPolicyResult properties."""
def test_should_log_property(self):
result = CommandPolicyResult(
decision=PolicyDecision.ALLOW,
original_command="ls",
masked_command="ls",
command_hash="abc123",
)
assert result.should_log is True
def test_is_blocked_property(self):
result = CommandPolicyResult(
decision=PolicyDecision.BLOCK,
original_command="cat /etc/shadow",
reason="Sensitive file access",
)
assert result.is_blocked is True
assert result.should_log is False
@@ -0,0 +1,295 @@
"""
Unit tests for Terminal Command History API endpoints.
Tests cover:
- GET /api/terminal/{host_id}/command-history
- GET /api/terminal/{host_id}/command-history/unique
- GET /api/terminal/command-history (global)
- DELETE /api/terminal/{host_id}/command-history
- POST /api/terminal/command-history/purge
"""
import pytest
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock, patch
from httpx import AsyncClient
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.terminal_command_log import TerminalCommandLog
class TestGetHostCommandHistory:
"""Test GET /api/terminal/{host_id}/command-history endpoint."""
@pytest.mark.asyncio
async def test_returns_command_history(self, client: AsyncClient, db_session: AsyncSession, auth_headers):
"""Test successful retrieval of command history."""
# Create test data
from app.crud.terminal_command_log import TerminalCommandLogRepository
from app.crud.host import HostRepository
# First ensure we have a host
host_repo = HostRepository(db_session)
hosts = await host_repo.list()
if hosts:
host_id = hosts[0].id
response = await client.get(
f"/api/terminal/{host_id}/command-history",
headers=auth_headers,
)
assert response.status_code == 200
data = response.json()
assert "commands" in data
assert "total" in data
assert isinstance(data["commands"], list)
@pytest.mark.asyncio
async def test_returns_404_for_unknown_host(self, client: AsyncClient, auth_headers):
"""Test 404 response for non-existent host."""
response = await client.get(
"/api/terminal/nonexistent-host-id/command-history",
headers=auth_headers,
)
assert response.status_code == 404
@pytest.mark.asyncio
async def test_filters_by_query(self, client: AsyncClient, db_session: AsyncSession, auth_headers):
"""Test search query filtering."""
from app.crud.host import HostRepository
host_repo = HostRepository(db_session)
hosts = await host_repo.list()
if hosts:
host_id = hosts[0].id
response = await client.get(
f"/api/terminal/{host_id}/command-history?query=ls",
headers=auth_headers,
)
assert response.status_code == 200
data = response.json()
assert data["query"] == "ls"
@pytest.mark.asyncio
async def test_respects_limit(self, client: AsyncClient, db_session: AsyncSession, auth_headers):
"""Test limit parameter."""
from app.crud.host import HostRepository
host_repo = HostRepository(db_session)
hosts = await host_repo.list()
if hosts:
host_id = hosts[0].id
response = await client.get(
f"/api/terminal/{host_id}/command-history?limit=10",
headers=auth_headers,
)
assert response.status_code == 200
data = response.json()
assert len(data["commands"]) <= 10
@pytest.mark.asyncio
async def test_caps_limit_at_100(self, client: AsyncClient, db_session: AsyncSession, auth_headers):
"""Test that limit is capped at 100."""
from app.crud.host import HostRepository
host_repo = HostRepository(db_session)
hosts = await host_repo.list()
if hosts:
host_id = hosts[0].id
response = await client.get(
f"/api/terminal/{host_id}/command-history?limit=500",
headers=auth_headers,
)
assert response.status_code == 200
# Should still work, just capped internally
class TestGetHostUniqueCommands:
"""Test GET /api/terminal/{host_id}/command-history/unique endpoint."""
@pytest.mark.asyncio
async def test_returns_unique_commands(self, client: AsyncClient, db_session: AsyncSession, auth_headers):
"""Test retrieval of unique commands."""
from app.crud.host import HostRepository
host_repo = HostRepository(db_session)
hosts = await host_repo.list()
if hosts:
host_id = hosts[0].id
response = await client.get(
f"/api/terminal/{host_id}/command-history/unique",
headers=auth_headers,
)
assert response.status_code == 200
data = response.json()
assert "commands" in data
assert "total" in data
assert "host_id" in data
class TestGetGlobalCommandHistory:
"""Test GET /api/terminal/command-history endpoint."""
@pytest.mark.asyncio
async def test_returns_global_history(self, client: AsyncClient, auth_headers):
"""Test retrieval of global command history."""
response = await client.get(
"/api/terminal/command-history",
headers=auth_headers,
)
assert response.status_code == 200
data = response.json()
assert "commands" in data
assert isinstance(data["commands"], list)
@pytest.mark.asyncio
async def test_filters_by_host_id(self, client: AsyncClient, db_session: AsyncSession, auth_headers):
"""Test filtering by host_id."""
from app.crud.host import HostRepository
host_repo = HostRepository(db_session)
hosts = await host_repo.list()
if hosts:
host_id = hosts[0].id
response = await client.get(
f"/api/terminal/command-history?host_id={host_id}",
headers=auth_headers,
)
assert response.status_code == 200
data = response.json()
assert data["host_id"] == host_id
class TestClearHostCommandHistory:
"""Test DELETE /api/terminal/{host_id}/command-history endpoint."""
@pytest.mark.asyncio
async def test_requires_admin_role(self, client: AsyncClient, db_session: AsyncSession, auth_headers):
"""Test that only admins can clear history."""
from app.crud.host import HostRepository
host_repo = HostRepository(db_session)
hosts = await host_repo.list()
if hosts:
host_id = hosts[0].id
# This test depends on the auth setup
# If the test user is admin, it should succeed
# If not admin, it should return 403
response = await client.delete(
f"/api/terminal/{host_id}/command-history",
headers=auth_headers,
)
# Either 200 (admin) or 403 (not admin) is acceptable
assert response.status_code in [200, 403]
@pytest.mark.asyncio
async def test_returns_404_for_unknown_host(self, client: AsyncClient, auth_headers):
"""Test 404 for non-existent host."""
response = await client.delete(
"/api/terminal/nonexistent-host-id/command-history",
headers=auth_headers,
)
# Either 404 (host not found) or 403 (not admin) is acceptable
assert response.status_code in [403, 404]
class TestPurgeCommandHistory:
"""Test POST /api/terminal/command-history/purge endpoint."""
@pytest.mark.asyncio
async def test_requires_admin_role(self, client: AsyncClient, auth_headers):
"""Test that only admins can purge history."""
response = await client.post(
"/api/terminal/command-history/purge?days=30",
headers=auth_headers,
)
# Either 200 (admin) or 403 (not admin) is acceptable
assert response.status_code in [200, 403]
@pytest.mark.asyncio
async def test_accepts_days_parameter(self, client: AsyncClient, auth_headers):
"""Test days parameter is accepted."""
response = await client.post(
"/api/terminal/command-history/purge?days=7",
headers=auth_headers,
)
# Either 200 (admin) or 403 (not admin) is acceptable
assert response.status_code in [200, 403]
class TestCommandHistorySchemas:
"""Test schema validation for command history responses."""
@pytest.mark.asyncio
async def test_command_history_item_schema(self, client: AsyncClient, db_session: AsyncSession, auth_headers):
"""Test that response matches CommandHistoryItem schema."""
from app.crud.host import HostRepository
host_repo = HostRepository(db_session)
hosts = await host_repo.list()
if hosts:
host_id = hosts[0].id
response = await client.get(
f"/api/terminal/{host_id}/command-history",
headers=auth_headers,
)
assert response.status_code == 200
data = response.json()
for cmd in data["commands"]:
assert "id" in cmd
assert "command" in cmd
assert "created_at" in cmd
@pytest.mark.asyncio
async def test_unique_command_item_schema(self, client: AsyncClient, db_session: AsyncSession, auth_headers):
"""Test that response matches UniqueCommandItem schema."""
from app.crud.host import HostRepository
host_repo = HostRepository(db_session)
hosts = await host_repo.list()
if hosts:
host_id = hosts[0].id
response = await client.get(
f"/api/terminal/{host_id}/command-history/unique",
headers=auth_headers,
)
assert response.status_code == 200
data = response.json()
for cmd in data["commands"]:
assert "command" in cmd
assert "command_hash" in cmd
assert "last_used" in cmd
assert "execution_count" in cmd
@@ -0,0 +1,371 @@
"""
Unit tests for the Terminal Command Logger.
Tests cover:
- Buffer management (character accumulation)
- Enter detection
- Backspace handling
- Ctrl+U handling
- Command validation integration
"""
import pytest
import asyncio
from unittest.mock import AsyncMock, MagicMock, patch
from app.services.terminal_command_logger import (
TerminalCommandLogger,
SessionContext,
get_command_logger,
)
from app.security.command_policy import CommandPolicy, PolicyDecision
class TestSessionContext:
"""Test SessionContext dataclass."""
def test_creates_session_context(self):
ctx = SessionContext(
session_id="test-session-123",
host_id="host-1",
host_name="test-host",
user_id="user-1",
username="testuser",
)
assert ctx.session_id == "test-session-123"
assert ctx.host_id == "host-1"
assert ctx.buffer == ""
assert ctx.commands_logged == 0
assert ctx.commands_blocked == 0
class TestTerminalCommandLoggerBuffer:
"""Test buffer management in TerminalCommandLogger."""
@pytest.fixture
def logger(self):
return TerminalCommandLogger()
@pytest.mark.asyncio
async def test_creates_session(self, logger):
ctx = await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
user_id="user-1",
username="testuser",
)
assert ctx is not None
assert ctx.session_id == "sess-1"
assert logger.get_session("sess-1") is ctx
@pytest.mark.asyncio
async def test_removes_session(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
removed = await logger.remove_session("sess-1")
assert removed is not None
assert logger.get_session("sess-1") is None
@pytest.mark.asyncio
async def test_accumulates_printable_chars(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
# Send printable characters
await logger.process_input("sess-1", b"ls ")
ctx = logger.get_session("sess-1")
assert ctx.buffer == "ls "
await logger.process_input("sess-1", b"-la")
assert ctx.buffer == "ls -la"
@pytest.mark.asyncio
async def test_handles_backspace(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
# Type and then backspace
await logger.process_input("sess-1", b"lss")
ctx = logger.get_session("sess-1")
assert ctx.buffer == "lss"
# Backspace (DEL character)
await logger.process_input("sess-1", b"\x7f")
assert ctx.buffer == "ls"
@pytest.mark.asyncio
async def test_handles_ctrl_u(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
# Type something
await logger.process_input("sess-1", b"some command")
ctx = logger.get_session("sess-1")
assert ctx.buffer == "some command"
# Ctrl+U clears line
await logger.process_input("sess-1", b"\x15")
assert ctx.buffer == ""
@pytest.mark.asyncio
async def test_handles_ctrl_c(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
await logger.process_input("sess-1", b"some command")
ctx = logger.get_session("sess-1")
# Ctrl+C clears buffer
await logger.process_input("sess-1", b"\x03")
assert ctx.buffer == ""
class TestTerminalCommandLoggerEnterDetection:
"""Test Enter key detection and command processing."""
@pytest.fixture
def logger(self):
return TerminalCommandLogger()
@pytest.mark.asyncio
async def test_detects_enter_cr(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
# Type command and press Enter (CR)
results = await logger.process_input("sess-1", b"ls -la\r")
# Should have processed the command
assert len(results) == 1
# Buffer should be cleared
ctx = logger.get_session("sess-1")
assert ctx.buffer == ""
@pytest.mark.asyncio
async def test_detects_enter_lf(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
# Type command and press Enter (LF)
results = await logger.process_input("sess-1", b"pwd\n")
assert len(results) == 1
ctx = logger.get_session("sess-1")
assert ctx.buffer == ""
@pytest.mark.asyncio
async def test_detects_enter_crlf(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
# Type command and press Enter (CRLF)
results = await logger.process_input("sess-1", b"whoami\r\n")
assert len(results) == 1
@pytest.mark.asyncio
async def test_ignores_empty_command(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
# Just press Enter with empty buffer
results = await logger.process_input("sess-1", b"\r")
assert len(results) == 0
@pytest.mark.asyncio
async def test_multiple_commands_in_stream(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
# Multiple commands in one stream
results = await logger.process_input("sess-1", b"ls\rpwd\r")
assert len(results) == 2
class TestTerminalCommandLoggerPolicyIntegration:
"""Test integration with CommandPolicy."""
@pytest.fixture
def logger(self):
return TerminalCommandLogger()
@pytest.mark.asyncio
async def test_logs_allowed_command(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
results = await logger.process_input("sess-1", b"ls -la\r")
assert len(results) == 1
assert results[0].should_log
ctx = logger.get_session("sess-1")
assert ctx.commands_logged == 1
@pytest.mark.asyncio
async def test_blocks_sensitive_command(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
results = await logger.process_input("sess-1", b"cat /etc/shadow\r")
assert len(results) == 1
assert results[0].is_blocked
ctx = logger.get_session("sess-1")
assert ctx.commands_blocked == 1
@pytest.mark.asyncio
async def test_callback_called_for_allowed(self, logger):
callback = AsyncMock()
logger.set_log_callback(callback)
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
user_id="user-1",
username="testuser",
)
await logger.process_input("sess-1", b"pwd\r")
callback.assert_called_once()
call_kwargs = callback.call_args.kwargs
assert call_kwargs["host_id"] == "host-1"
assert call_kwargs["command"] == "pwd"
@pytest.mark.asyncio
async def test_callback_called_for_blocked(self, logger):
callback = AsyncMock()
logger.set_log_callback(callback)
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
await logger.process_input("sess-1", b"cat ~/.ssh/id_rsa\r")
callback.assert_called_once()
call_kwargs = callback.call_args.kwargs
assert call_kwargs["is_blocked"] is True
class TestTerminalCommandLoggerStats:
"""Test statistics tracking."""
@pytest.fixture
def logger(self):
return TerminalCommandLogger()
@pytest.mark.asyncio
async def test_get_stats(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
await logger.process_input("sess-1", b"ls\r")
await logger.process_input("sess-1", b"pwd\r")
stats = logger.get_stats()
assert stats["active_sessions"] == 1
assert stats["total_commands_logged"] >= 0
class TestFlushBuffer:
"""Test buffer flushing."""
@pytest.fixture
def logger(self):
return TerminalCommandLogger()
@pytest.mark.asyncio
async def test_flush_buffer_with_content(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
# Add content to buffer without Enter
await logger.process_input("sess-1", b"incomplete command")
# Flush buffer
result = await logger.flush_buffer("sess-1")
# Should process the command
assert result is not None
# Buffer should be cleared
ctx = logger.get_session("sess-1")
assert ctx.buffer == ""
@pytest.mark.asyncio
async def test_flush_empty_buffer(self, logger):
await logger.create_session(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
)
# Flush empty buffer
result = await logger.flush_buffer("sess-1")
assert result is None
@pytest.mark.asyncio
async def test_flush_nonexistent_session(self, logger):
result = await logger.flush_buffer("nonexistent")
assert result is None
class TestGlobalInstance:
"""Test global instance management."""
def test_get_command_logger_returns_singleton(self):
logger1 = get_command_logger()
logger2 = get_command_logger()
assert logger1 is logger2
+563
View File
@@ -0,0 +1,563 @@
"""
Tests for terminal session management.
Tests cover:
- Session reuse for same user/host/mode
- Session limit with rich error response
- Idempotent close
- Heartbeat updates last_seen_at
- GC cleans up expired/idle sessions
- Close beacon endpoint
"""
import pytest
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock, patch
from app.models.terminal_session import (
TerminalSession,
SESSION_STATUS_ACTIVE,
SESSION_STATUS_CLOSED,
SESSION_STATUS_EXPIRED,
CLOSE_REASON_USER,
CLOSE_REASON_TTL,
CLOSE_REASON_IDLE,
CLOSE_REASON_CLIENT_LOST,
)
from app.crud.terminal_session import TerminalSessionRepository
from app.schemas.terminal import (
SessionLimitError,
HeartbeatResponse,
ActiveSessionInfo,
)
# ============================================================================
# Fixtures
# ============================================================================
@pytest.fixture
def mock_session():
"""Create a mock terminal session."""
now = datetime.now(timezone.utc)
return TerminalSession(
id="test-session-123",
host_id="host-1",
host_name="test-host",
host_ip="192.168.1.100",
user_id="user-1",
username="testuser",
token_hash="hash123",
ttyd_port=7680,
ttyd_pid=12345,
mode="embedded",
status=SESSION_STATUS_ACTIVE,
created_at=now,
last_seen_at=now,
expires_at=now + timedelta(minutes=30),
)
@pytest.fixture
def mock_db_session():
"""Create a mock database session."""
session = AsyncMock()
session.execute = AsyncMock()
session.flush = AsyncMock()
session.commit = AsyncMock()
return session
# ============================================================================
# Session Reuse Tests
# ============================================================================
class TestSessionReuse:
"""Tests for session reuse functionality."""
@pytest.mark.asyncio
async def test_find_reusable_session_returns_matching_session(self, mock_db_session, mock_session):
"""Should find and return an existing active session for same user/host/mode."""
# Setup mock to return the session
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = mock_session
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
result = await repo.find_reusable_session(
user_id="user-1",
host_id="host-1",
mode="embedded",
idle_timeout_seconds=120
)
assert result is not None
assert result.id == mock_session.id
assert result.host_id == "host-1"
assert result.user_id == "user-1"
@pytest.mark.asyncio
async def test_find_reusable_session_returns_none_for_different_mode(self, mock_db_session):
"""Should not find session if mode is different."""
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = None
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
result = await repo.find_reusable_session(
user_id="user-1",
host_id="host-1",
mode="popout", # Different mode
idle_timeout_seconds=120
)
assert result is None
@pytest.mark.asyncio
async def test_find_reusable_session_excludes_idle_sessions(self, mock_db_session):
"""Should not return sessions that are idle (last_seen too old)."""
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = None
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
# With a very short idle timeout, session should be considered idle
result = await repo.find_reusable_session(
user_id="user-1",
host_id="host-1",
mode="embedded",
idle_timeout_seconds=1 # Very short timeout
)
assert result is None
# ============================================================================
# Session Limit Tests
# ============================================================================
class TestSessionLimit:
"""Tests for session limit handling."""
def test_session_limit_error_schema(self):
"""SessionLimitError schema should contain all required fields."""
error = SessionLimitError(
message="Maximum sessions reached",
max_active=3,
current_count=3,
active_sessions=[
ActiveSessionInfo(
session_id="sess-1",
host_id="host-1",
host_name="test-host",
mode="embedded",
age_seconds=300,
last_seen_seconds=10,
)
],
suggested_actions=["close_oldest", "close_session:sess-1"],
can_reuse=False,
)
assert error.error == "SESSION_LIMIT"
assert error.max_active == 3
assert error.current_count == 3
assert len(error.active_sessions) == 1
assert "close_oldest" in error.suggested_actions
def test_session_limit_error_with_reusable_session(self):
"""SessionLimitError should indicate when a session can be reused."""
error = SessionLimitError(
message="Maximum sessions reached",
max_active=3,
current_count=3,
active_sessions=[],
suggested_actions=["reuse_existing", "close_oldest"],
can_reuse=True,
reusable_session_id="sess-reusable",
)
assert error.can_reuse is True
assert error.reusable_session_id == "sess-reusable"
assert "reuse_existing" in error.suggested_actions
# ============================================================================
# Idempotent Close Tests
# ============================================================================
class TestIdempotentClose:
"""Tests for idempotent session close."""
@pytest.mark.asyncio
async def test_close_already_closed_session(self, mock_db_session, mock_session):
"""Closing an already closed session should succeed (idempotent)."""
mock_session.status = SESSION_STATUS_CLOSED
mock_session.closed_at = datetime.now(timezone.utc)
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = mock_session
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
# Close should not raise, even if already closed
result = await repo.close_session(mock_session.id)
# Status should remain closed
assert result.status == SESSION_STATUS_CLOSED
@pytest.mark.asyncio
async def test_close_nonexistent_session(self, mock_db_session):
"""Closing a nonexistent session should return None."""
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = None
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
result = await repo.close_session("nonexistent-session")
assert result is None
# ============================================================================
# Heartbeat Tests
# ============================================================================
class TestHeartbeat:
"""Tests for heartbeat functionality."""
@pytest.mark.asyncio
async def test_heartbeat_updates_last_seen(self, mock_db_session, mock_session):
"""Heartbeat should update last_seen_at timestamp."""
original_last_seen = mock_session.last_seen_at
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = mock_session
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
result = await repo.update_last_seen(mock_session.id)
assert result is not None
assert result.last_seen_at >= original_last_seen
@pytest.mark.asyncio
async def test_heartbeat_on_closed_session_does_nothing(self, mock_db_session, mock_session):
"""Heartbeat on closed session should not update last_seen."""
mock_session.status = SESSION_STATUS_CLOSED
original_last_seen = mock_session.last_seen_at
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = mock_session
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
result = await repo.update_last_seen(mock_session.id)
# last_seen should not be updated for closed sessions
assert result.last_seen_at == original_last_seen
def test_heartbeat_response_schema(self):
"""HeartbeatResponse schema should contain all required fields."""
response = HeartbeatResponse(
session_id="test-session",
status=SESSION_STATUS_ACTIVE,
last_seen_at=datetime.now(timezone.utc),
remaining_seconds=1500,
healthy=True,
)
assert response.session_id == "test-session"
assert response.status == SESSION_STATUS_ACTIVE
assert response.healthy is True
assert response.remaining_seconds == 1500
# ============================================================================
# GC (Garbage Collection) Tests
# ============================================================================
class TestGarbageCollection:
"""Tests for session garbage collection."""
@pytest.mark.asyncio
async def test_list_expired_sessions(self, mock_db_session):
"""Should list sessions past their TTL."""
now = datetime.now(timezone.utc)
expired_session = TerminalSession(
id="expired-1",
host_id="host-1",
host_name="test-host",
host_ip="192.168.1.100",
user_id="user-1",
token_hash="hash",
ttyd_port=7680,
mode="embedded",
status=SESSION_STATUS_ACTIVE,
created_at=now - timedelta(hours=1),
last_seen_at=now - timedelta(minutes=5),
expires_at=now - timedelta(minutes=10), # Expired
)
mock_result = MagicMock()
mock_result.scalars.return_value.all.return_value = [expired_session]
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
expired = await repo.list_expired()
assert len(expired) == 1
assert expired[0].id == "expired-1"
@pytest.mark.asyncio
async def test_list_idle_sessions(self, mock_db_session):
"""Should list sessions without recent heartbeat."""
now = datetime.now(timezone.utc)
idle_session = TerminalSession(
id="idle-1",
host_id="host-1",
host_name="test-host",
host_ip="192.168.1.100",
user_id="user-1",
token_hash="hash",
ttyd_port=7680,
mode="embedded",
status=SESSION_STATUS_ACTIVE,
created_at=now - timedelta(hours=1),
last_seen_at=now - timedelta(minutes=5), # Idle for 5 minutes
expires_at=now + timedelta(minutes=25), # Not expired
)
mock_result = MagicMock()
mock_result.scalars.return_value.all.return_value = [idle_session]
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
# With 120 second idle timeout, 5 minutes idle should be detected
idle = await repo.list_idle(idle_timeout_seconds=120)
assert len(idle) == 1
assert idle[0].id == "idle-1"
@pytest.mark.asyncio
async def test_list_stale_sessions_combines_expired_and_idle(self, mock_db_session):
"""Should list both expired and idle sessions."""
now = datetime.now(timezone.utc)
mock_result = MagicMock()
mock_result.scalars.return_value.all.return_value = []
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
stale = await repo.list_stale_sessions(
ttl_seconds=1800,
idle_timeout_seconds=120
)
# Query should be executed
mock_db_session.execute.assert_called()
@pytest.mark.asyncio
async def test_mark_expired_sets_correct_reason(self, mock_db_session, mock_session):
"""mark_expired should set reason_closed to TTL."""
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = mock_session
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
result = await repo.mark_expired(mock_session.id)
assert result.status == SESSION_STATUS_EXPIRED
assert result.reason_closed == CLOSE_REASON_TTL
assert result.closed_at is not None
@pytest.mark.asyncio
async def test_mark_idle_sets_correct_reason(self, mock_db_session, mock_session):
"""mark_idle should set reason_closed to IDLE."""
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = mock_session
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
result = await repo.mark_idle(mock_session.id)
assert result.status == SESSION_STATUS_EXPIRED
assert result.reason_closed == CLOSE_REASON_IDLE
assert result.closed_at is not None
# ============================================================================
# Close Beacon Tests
# ============================================================================
class TestCloseBeacon:
"""Tests for close beacon endpoint."""
@pytest.mark.asyncio
async def test_mark_client_lost_sets_correct_reason(self, mock_db_session, mock_session):
"""mark_client_lost should set reason_closed to CLIENT_LOST."""
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = mock_session
mock_db_session.execute.return_value = mock_result
repo = TerminalSessionRepository(mock_db_session)
result = await repo.mark_client_lost(mock_session.id)
assert result.status == SESSION_STATUS_CLOSED
assert result.reason_closed == CLOSE_REASON_CLIENT_LOST
assert result.closed_at is not None
# ============================================================================
# Terminal Service Tests
# ============================================================================
class TestTerminalService:
"""Tests for terminal service functionality."""
def test_metrics_tracking(self):
"""Service should track metrics for observability."""
from app.services.terminal_service import TerminalService
service = TerminalService()
# Record some metrics
service.record_session_created(reused=False)
service.record_session_created(reused=True)
service.record_session_limit_hit()
metrics = service.get_metrics()
assert metrics["sessions_created"] == 1
assert metrics["sessions_reused"] == 1
assert metrics["session_limit_hits"] == 1
assert "active_processes" in metrics
assert "gc_running" in metrics
def test_token_generation_and_verification(self):
"""Service should generate and verify tokens correctly."""
from app.services.terminal_service import TerminalService
service = TerminalService()
token, token_hash = service.generate_session_token()
assert token is not None
assert token_hash is not None
assert len(token) > 32
assert len(token_hash) == 64 # SHA256 hex
# Verify correct token
assert service.verify_token(token, token_hash) is True
# Verify incorrect token
assert service.verify_token("wrong-token", token_hash) is False
def test_session_id_generation(self):
"""Service should generate unique session IDs."""
from app.services.terminal_service import TerminalService
service = TerminalService()
id1 = service.generate_session_id()
id2 = service.generate_session_id()
assert id1 != id2
assert len(id1) == 64 # 32 bytes hex
def test_to_utc_aware_handles_naive_datetime(self):
"""GC should not crash when SQLite returns naive datetimes."""
from app.services.terminal_service import TerminalService
service = TerminalService()
naive = datetime(2025, 1, 1, 12, 0, 0)
aware = service._to_utc_aware(naive)
assert aware is not None
assert aware.tzinfo is not None
assert aware.utcoffset() == timedelta(0)
@pytest.mark.asyncio
async def test_gc_cycle_expired_comparison_with_naive_expires_at(self, mock_db_session):
"""_run_gc_cycle should not raise when expires_at is naive."""
from app.services.terminal_service import TerminalService
from app.crud import terminal_session as terminal_session_module
service = TerminalService()
now = datetime.now(timezone.utc)
session = TerminalSession(
id="naive-exp-1",
host_id="host-1",
host_name="test-host",
host_ip="192.168.1.100",
user_id="user-1",
token_hash="hash",
ttyd_port=7680,
mode="embedded",
status=SESSION_STATUS_ACTIVE,
created_at=now - timedelta(minutes=10),
last_seen_at=now - timedelta(minutes=1),
expires_at=datetime.now() - timedelta(minutes=1),
)
# Fake repo that returns a session with naive expires_at
fake_repo = AsyncMock()
fake_repo.list_stale_sessions.return_value = [session]
fake_repo.mark_expired.return_value = session
fake_repo.mark_idle.return_value = session
# Fake async session factory
class _Factory:
async def __aenter__(self_inner):
return mock_db_session
async def __aexit__(self_inner, exc_type, exc, tb):
return False
service._db_session_factory = lambda: _Factory()
with patch.object(terminal_session_module, "TerminalSessionRepository", return_value=fake_repo):
service.terminate_session = AsyncMock()
service.release_port = AsyncMock()
await service._run_gc_cycle()
# ============================================================================
# Model Tests
# ============================================================================
class TestTerminalSessionModel:
"""Tests for TerminalSession model."""
def test_status_constants(self):
"""Status constants should be defined correctly."""
assert SESSION_STATUS_ACTIVE == "active"
assert SESSION_STATUS_CLOSED == "closed"
assert SESSION_STATUS_EXPIRED == "expired"
def test_close_reason_constants(self):
"""Close reason constants should be defined correctly."""
assert CLOSE_REASON_USER == "user_close"
assert CLOSE_REASON_TTL == "ttl"
assert CLOSE_REASON_IDLE == "idle"
assert CLOSE_REASON_CLIENT_LOST == "client_lost"
def test_session_repr(self, mock_session):
"""Session repr should be informative."""
repr_str = repr(mock_session)
assert "TerminalSession" in repr_str
assert mock_session.host_name in repr_str
assert mock_session.status in repr_str