Files

618 lines
23 KiB
Python
Raw Permalink Normal View History

"""Tests for the External Agent API — route layer.
Uses AsyncMock for all DB operations, stream_chat, and rate limiting.
No real DB or LLM connections required.
"""
from __future__ import annotations
import json
import uuid
from datetime import UTC, datetime
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import FastAPI
from httpx import ASGITransport, AsyncClient
from app.plugins.builtins.ai_assistant.external_api import router as external_api_router
from app.plugins.builtins.ai_assistant.models import AIAgent
# Override conftest DB fixtures — these tests use mocks, no real DB needed
@pytest.fixture(autouse=True, scope="session")
def db_setup():
"""No-op override of conftest db_setup."""
yield
@pytest.fixture(autouse=True)
def clean_tables(db_setup):
"""No-op override of conftest clean_tables."""
yield
# ─── Helpers ───
def _make_agent(
*,
tenant_id: uuid.UUID | None = None,
name: str = "Test Agent",
is_active: bool = True,
tool_ids: list[str] | None = None,
) -> AIAgent:
"""Create an AIAgent instance with defaults."""
return AIAgent(
id=uuid.uuid4(),
tenant_id=tenant_id or uuid.uuid4(),
name=name,
description="A test agent",
system_prompt="You are a helpful assistant.",
tool_ids=tool_ids or [],
is_active=is_active,
is_default=False,
config={},
created_at=datetime.now(UTC),
updated_at=datetime.now(UTC),
)
def _mock_session() -> AsyncMock:
"""Create a mock AsyncSession."""
session = AsyncMock()
session.add = MagicMock()
session.flush = AsyncMock()
session.execute = AsyncMock()
session.commit = AsyncMock()
session.rollback = AsyncMock()
return session
def _create_external_api_app(
*,
db: AsyncMock | None = None,
current_user: dict | None = None,
) -> FastAPI:
"""Create a minimal FastAPI app with external_api router and mocked dependencies."""
app = FastAPI()
app.include_router(external_api_router)
if db is None:
db = _mock_session()
if current_user is None:
current_user = {
"user_id": str(uuid.uuid4()),
"tenant_id": str(uuid.uuid4()),
"is_system_admin": True,
"role": "admin",
"permissions": [],
"denied_permissions": [],
"field_permissions": {},
"token_prefix": "test_token",
}
async def _mock_get_db():
yield db
async def _mock_get_current_user_bearer():
return current_user
async def _mock_set_tenant_context(session, tenant_id):
pass
async def _mock_check_rate_limit(redis_key, max_attempts, window_seconds):
pass
from app.deps import get_current_user_bearer, get_current_user
from app.core.db import get_db, set_tenant_context
app.dependency_overrides[get_db] = _mock_get_db
app.dependency_overrides[get_current_user_bearer] = _mock_get_current_user_bearer
app.dependency_overrides[get_current_user] = _mock_get_current_user_bearer
# Patch set_tenant_context and check_rate_limit at module level
return app
# ─── Run Agent Tests ───
class TestRunAgentExternal:
"""Tests for POST /api/v1/external/agent/{agent_id}/run."""
@pytest.mark.asyncio
async def test_run_agent_success(self):
"""POST /run executes agent and returns response."""
tenant_id = uuid.uuid4()
agent = _make_agent(tenant_id=tenant_id, is_active=True)
db = _mock_session()
# First execute: find agent
agent_result = MagicMock()
agent_result.scalar_one_or_none.return_value = agent
db.execute.return_value = agent_result
current_user = {
"user_id": str(uuid.uuid4()),
"tenant_id": str(tenant_id),
"is_system_admin": True,
"role": "admin",
"permissions": [],
"denied_permissions": [],
"field_permissions": {},
"token_prefix": "test_token",
}
app = _create_external_api_app(db=db, current_user=current_user)
# Mock stream_chat to return chunks
async def _mock_stream_chat(*args, **kwargs):
yield 'data: {"content": "Hello "}\n\n'
yield 'data: {"content": "world!"}\n\n'
yield "data: [DONE]\n\n"
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api.get_db") as mock_get_db,
patch("app.plugins.builtins.ai_assistant.services.stream_chat", _mock_stream_chat),
):
# Override get_db to return an async context manager for the inner stream_db
class _FakeAsyncCtxMgr:
async def __aenter__(self):
return db
async def __aexit__(self, *args):
pass
mock_get_db.return_value = _FakeAsyncCtxMgr()
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
f"/api/v1/external/agent/{agent.id}/run",
json={"message": "Say hello"},
)
assert resp.status_code == 200
data = resp.json()
assert "response" in data
assert data["agent_id"] == str(agent.id)
assert "session_id" in data
assert data["tokens_used"] > 0
@pytest.mark.asyncio
async def test_run_agent_not_found(self):
"""POST /run returns 404 when agent does not exist."""
db = _mock_session()
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = None
db.execute.return_value = mock_result
app = _create_external_api_app(db=db)
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
f"/api/v1/external/agent/{uuid.uuid4()}/run",
json={"message": "test"},
)
assert resp.status_code == 404
assert resp.json()["detail"] == "Agent not found"
@pytest.mark.asyncio
async def test_run_agent_inactive(self):
"""POST /run returns 400 when agent is inactive."""
agent = _make_agent(is_active=False)
db = _mock_session()
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = agent
db.execute.return_value = mock_result
app = _create_external_api_app(db=db)
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
f"/api/v1/external/agent/{agent.id}/run",
json={"message": "test"},
)
assert resp.status_code == 400
assert resp.json()["detail"] == "Agent is not active"
@pytest.mark.asyncio
async def test_run_agent_invalid_uuid(self):
"""POST /run returns 400 for invalid agent UUID."""
app = _create_external_api_app()
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
"/api/v1/external/agent/not-a-uuid/run",
json={"message": "test"},
)
assert resp.status_code == 400
assert resp.json()["detail"] == "Invalid agent ID"
@pytest.mark.asyncio
async def test_run_agent_rate_limited(self):
"""POST /run returns 429 when rate limit is exceeded."""
from fastapi import HTTPException, status
agent = _make_agent(is_active=True)
db = _mock_session()
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = agent
db.execute.return_value = mock_result
app = _create_external_api_app(db=db)
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch(
"app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit",
AsyncMock(side_effect=HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="Rate limit exceeded",
)),
),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
f"/api/v1/external/agent/{agent.id}/run",
json={"message": "test"},
)
assert resp.status_code == 429
# ─── Get Agent Status Tests ───
class TestGetAgentStatusExternal:
"""Tests for GET /api/v1/external/agent/{agent_id}/status."""
@pytest.mark.asyncio
async def test_get_status_success(self):
"""GET /status returns agent status information."""
tenant_id = uuid.uuid4()
agent = _make_agent(tenant_id=tenant_id, name="Status Agent", tool_ids=["tool1", "tool2"])
db = _mock_session()
# First execute: find agent
agent_result = MagicMock()
agent_result.scalar_one_or_none.return_value = agent
# Second execute: count runs
count_result = MagicMock()
count_result.scalar.return_value = 5
db.execute.side_effect = [agent_result, count_result]
app = _create_external_api_app(db=db)
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.get(f"/api/v1/external/agent/{agent.id}/status")
assert resp.status_code == 200
data = resp.json()
assert data["agent_id"] == str(agent.id)
assert data["name"] == "Status Agent"
assert data["is_active"] is True
assert data["tool_count"] == 2
assert data["total_runs"] == 5
@pytest.mark.asyncio
async def test_get_status_not_found(self):
"""GET /status returns 404 when agent does not exist."""
db = _mock_session()
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = None
db.execute.return_value = mock_result
app = _create_external_api_app(db=db)
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.get(f"/api/v1/external/agent/{uuid.uuid4()}/status")
assert resp.status_code == 404
@pytest.mark.asyncio
async def test_get_status_invalid_uuid(self):
"""GET /status returns 400 for invalid agent UUID."""
app = _create_external_api_app()
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.get("/api/v1/external/agent/not-a-uuid/status")
assert resp.status_code == 400
assert resp.json()["detail"] == "Invalid agent ID"
@pytest.mark.asyncio
async def test_get_status_rate_limited(self):
"""GET /status returns 429 when rate limit is exceeded."""
from fastapi import HTTPException, status
app = _create_external_api_app()
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch(
"app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit",
AsyncMock(side_effect=HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="Rate limit exceeded",
)),
),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.get(f"/api/v1/external/agent/{uuid.uuid4()}/status")
assert resp.status_code == 429
# ─── Stream Agent Tests ───
class TestStreamAgentExternal:
"""Tests for POST /api/v1/external/agent/{agent_id}/stream."""
@pytest.mark.asyncio
async def test_stream_agent_success(self):
"""POST /stream returns SSE streaming response."""
tenant_id = uuid.uuid4()
agent = _make_agent(tenant_id=tenant_id, is_active=True)
db = _mock_session()
agent_result = MagicMock()
agent_result.scalar_one_or_none.return_value = agent
db.execute.return_value = agent_result
app = _create_external_api_app(db=db)
async def _mock_stream_chat(*args, **kwargs):
yield 'data: {"content": "Hello "}\n\n'
yield 'data: {"content": "world!"}\n\n'
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
patch("app.core.db.get_session_factory") as mock_factory,
patch("app.plugins.builtins.ai_assistant.services.stream_chat", _mock_stream_chat),
):
# Mock session factory for the streaming inner DB session
mock_session_ctx = AsyncMock()
mock_session_ctx.__aenter__ = AsyncMock(return_value=db)
mock_session_ctx.__aexit__ = AsyncMock(return_value=None)
mock_factory.return_value.return_value = mock_session_ctx
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
f"/api/v1/external/agent/{agent.id}/stream",
json={"message": "Say hello"},
)
assert resp.status_code == 200
assert "text/event-stream" in resp.headers.get("content-type", "")
# Response should contain SSE data
body = resp.text
assert "data:" in body
@pytest.mark.asyncio
async def test_stream_agent_not_found(self):
"""POST /stream returns 404 when agent does not exist."""
db = _mock_session()
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = None
db.execute.return_value = mock_result
app = _create_external_api_app(db=db)
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
f"/api/v1/external/agent/{uuid.uuid4()}/stream",
json={"message": "test"},
)
assert resp.status_code == 404
@pytest.mark.asyncio
async def test_stream_agent_inactive(self):
"""POST /stream returns 400 when agent is inactive."""
agent = _make_agent(is_active=False)
db = _mock_session()
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = agent
db.execute.return_value = mock_result
app = _create_external_api_app(db=db)
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
f"/api/v1/external/agent/{agent.id}/stream",
json={"message": "test"},
)
assert resp.status_code == 400
assert resp.json()["detail"] == "Agent is not active"
@pytest.mark.asyncio
async def test_stream_agent_invalid_uuid(self):
"""POST /stream returns 400 for invalid agent UUID."""
app = _create_external_api_app()
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
"/api/v1/external/agent/not-a-uuid/stream",
json={"message": "test"},
)
assert resp.status_code == 400
assert resp.json()["detail"] == "Invalid agent ID"
@pytest.mark.asyncio
async def test_stream_agent_rate_limited(self):
"""POST /stream returns 429 when rate limit is exceeded."""
from fastapi import HTTPException, status
app = _create_external_api_app()
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch(
"app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit",
AsyncMock(side_effect=HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="Rate limit exceeded",
)),
),
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
f"/api/v1/external/agent/{uuid.uuid4()}/stream",
json={"message": "test"},
)
assert resp.status_code == 429
# ─── Auth Tests ───
class TestExternalAgentAuth:
"""Tests for Bearer token authentication on external agent API."""
@pytest.mark.asyncio
async def test_bearer_token_required(self):
"""Endpoints require Bearer token authentication (mocked via dependency override)."""
# When get_current_user_bearer is not overridden, the request should fail
app = FastAPI()
app.include_router(external_api_router)
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
f"/api/v1/external/agent/{uuid.uuid4()}/run",
json={"message": "test"},
)
# Should return 401 or 403 since no auth is provided
assert resp.status_code in (401, 403, 422)
@pytest.mark.asyncio
async def test_user_context_passed_correctly(self):
"""The current_user dict from bearer auth is passed to stream_chat."""
tenant_id = uuid.uuid4()
user_id = uuid.uuid4()
agent = _make_agent(tenant_id=tenant_id, is_active=True)
db = _mock_session()
agent_result = MagicMock()
agent_result.scalar_one_or_none.return_value = agent
db.execute.return_value = agent_result
current_user = {
"user_id": str(user_id),
"tenant_id": str(tenant_id),
"is_system_admin": False,
"role": "editor",
"permissions": ["ai:write"],
"denied_permissions": [],
"field_permissions": {"annual_revenue": "read"},
"token_prefix": "abc123",
}
app = _create_external_api_app(db=db, current_user=current_user)
captured_context = {}
async def _mock_stream_chat(stream_db, session, agent_obj, message, user_context, tid):
captured_context.update(user_context)
yield 'data: {"content": "response"}\n\n'
yield "data: [DONE]\n\n"
with (
patch("app.plugins.builtins.ai_assistant.external_api.set_tenant_context", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api._check_external_rate_limit", AsyncMock()),
patch("app.plugins.builtins.ai_assistant.external_api.get_db") as mock_get_db,
patch("app.plugins.builtins.ai_assistant.services.stream_chat", _mock_stream_chat),
):
class _FakeAsyncCtxMgr:
async def __aenter__(self):
return db
async def __aexit__(self, *args):
pass
mock_get_db.return_value = _FakeAsyncCtxMgr()
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
resp = await client.post(
f"/api/v1/external/agent/{agent.id}/run",
json={"message": "test"},
)
assert resp.status_code == 200
# Verify user context was passed correctly
assert captured_context.get("user_id") == str(user_id)
assert captured_context.get("tenant_id") == str(tenant_id)
assert captured_context.get("role") == "editor"
assert "ai:write" in captured_context.get("permissions", [])