618 lines
23 KiB
Python
618 lines
23 KiB
Python
|
|
"""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", [])
|