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
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:
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user