This commit is contained in:
@@ -0,0 +1,232 @@
|
||||
"""Shared pytest fixtures for agent-platform test suite.
|
||||
|
||||
The agent_platform package uses flat module imports (e.g. ``from api import app``)
|
||||
configured by ``agent_platform/pyproject.toml``'s ``py-modules``. We mirror that
|
||||
by adding ``agent_platform/`` to ``sys.path`` so the same import paths work in
|
||||
tests without needing the package to be installed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import AsyncIterator
|
||||
|
||||
import aiosqlite
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
# Make ``agent_platform/`` flat modules importable as if the package was installed.
|
||||
ROOT = Path(__file__).resolve().parent.parent
|
||||
AGENT_PLATFORM_DIR = ROOT / "agent_platform"
|
||||
if str(AGENT_PLATFORM_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(AGENT_PLATFORM_DIR))
|
||||
|
||||
# Set AUTH_TOKEN env var BEFORE importing config so the settings instance
|
||||
# sees a strong token (the production boot-guard requires >= 32 chars).
|
||||
# Individual fixtures monkeypatch settings.AUTH_TOKEN back to the test value.
|
||||
os.environ.setdefault("AUTH_TOKEN", "x" * 64)
|
||||
os.environ.setdefault("DB_PATH", str(ROOT / "tests" / ".unused.db"))
|
||||
os.environ.setdefault("LLM_API_KEY", "")
|
||||
os.environ.setdefault("MCP_SERVER_URL", "http://localhost:8501/mcp")
|
||||
|
||||
from config import settings # noqa: E402 (must come after sys.path/env setup)
|
||||
import config as config_module # noqa: E402
|
||||
import api as api_module # noqa: E402
|
||||
import auth as auth_module # noqa: E402
|
||||
import db as db_module # noqa: E402
|
||||
import llm as llm_module # noqa: E402
|
||||
import mcp_client as mcp_client_module # noqa: E402
|
||||
import agent as agent_module # noqa: E402
|
||||
from api import app # noqa: E402
|
||||
|
||||
|
||||
# === Paths & Auth ===
|
||||
|
||||
AUTH_TOKEN_VALUE = "change-me-in-production"
|
||||
AUTH_HEADERS = {"Authorization": f"Bearer {AUTH_TOKEN_VALUE}"}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tmp_db_path(tmp_path) -> str:
|
||||
"""Return a unique DB path under pytest's tmp_path."""
|
||||
return str(tmp_path / "test.db")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def auth_headers() -> dict:
|
||||
"""Authorization header dict using the literal default AUTH_TOKEN."""
|
||||
return dict(AUTH_HEADERS)
|
||||
|
||||
|
||||
# === DB + Settings ===
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def clean_db(tmp_db_path, monkeypatch) -> AsyncIterator[aiosqlite.Connection]:
|
||||
"""Point settings at tmp_db_path, init schema, yield a live connection."""
|
||||
monkeypatch.setattr(config_module.settings, "DB_PATH", tmp_db_path)
|
||||
monkeypatch.setattr(config_module.settings, "AUTH_TOKEN", AUTH_TOKEN_VALUE)
|
||||
|
||||
await db_module.init_db()
|
||||
db = await aiosqlite.connect(tmp_db_path)
|
||||
try:
|
||||
yield db
|
||||
finally:
|
||||
await db.close()
|
||||
|
||||
|
||||
# === MCP Stub ===
|
||||
|
||||
class _StubMCPClient:
|
||||
"""In-memory MCP replacement used by API + agent layer."""
|
||||
|
||||
def __init__(self, tools=("echo", "calculate"), healthy: bool = True,
|
||||
health_exc: BaseException | None = None,
|
||||
connect_exc: BaseException | None = None,
|
||||
list_tools_exc: BaseException | None = None):
|
||||
self._tools = list(tools)
|
||||
self._healthy = healthy
|
||||
self._health_exc = health_exc
|
||||
self._connect_exc = connect_exc
|
||||
self._list_tools_exc = list_tools_exc
|
||||
|
||||
async def list_tools(self):
|
||||
if self._list_tools_exc:
|
||||
raise self._list_tools_exc
|
||||
return [{"name": t, "description": f"stub:{t}", "inputSchema": {"type": "object", "properties": {}}} for t in self._tools]
|
||||
|
||||
async def health(self) -> bool:
|
||||
if self._health_exc:
|
||||
raise self._health_exc
|
||||
if self._connect_exc:
|
||||
raise self._connect_exc
|
||||
return self._healthy
|
||||
|
||||
async def connect(self):
|
||||
if self._connect_exc:
|
||||
raise self._connect_exc
|
||||
return None
|
||||
|
||||
async def disconnect(self):
|
||||
return None
|
||||
|
||||
async def call_tool(self, name, arguments):
|
||||
return {"name": name, "result": f"stub-{name}"}
|
||||
|
||||
async def get_tools_for_llm(self, allowed=None):
|
||||
tools = await self.list_tools()
|
||||
result = []
|
||||
for t in tools:
|
||||
name = t.get("name")
|
||||
if allowed and name not in allowed:
|
||||
continue
|
||||
result.append({
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": name,
|
||||
"description": t.get("description", ""),
|
||||
"parameters": t.get("inputSchema", {"type": "object", "properties": {}}),
|
||||
},
|
||||
})
|
||||
return result
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def stub_mcp():
|
||||
"""Default healthy stub returning ['echo', 'calculate']."""
|
||||
return _StubMCPClient()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def stub_mcp_unreachable():
|
||||
"""Stub whose health() raises CancelledError (S1 regression fixture)."""
|
||||
import asyncio
|
||||
return _StubMCPClient(connect_exc=asyncio.CancelledError())
|
||||
|
||||
|
||||
def AsyncMock_dummy():
|
||||
"""Async no-op function to replace close_mcp_client during tests."""
|
||||
async def _noop():
|
||||
return None
|
||||
return _noop
|
||||
|
||||
|
||||
# === HTTP Client ===
|
||||
|
||||
async def _override_get_db() -> AsyncIterator[aiosqlite.Connection]:
|
||||
"""Yield a connection to settings.DB_PATH for the duration of one request."""
|
||||
db = await aiosqlite.connect(config_module.settings.DB_PATH)
|
||||
try:
|
||||
yield db
|
||||
finally:
|
||||
await db.close()
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def app_client(clean_db, monkeypatch, stub_mcp) -> AsyncIterator:
|
||||
"""httpx.AsyncClient against the FastAPI app via ASGI transport.
|
||||
|
||||
* Uses ASGITransport WITHOUT lifespan events (boot guard would otherwise
|
||||
require a 32+ char AUTH_TOKEN, which is set via clean_db).
|
||||
* Overrides ``api.get_db`` so every request gets the tmp DB connection.
|
||||
* Stubs ``get_mcp_client()`` (in api + agent modules, not just mcp_client)
|
||||
so list_tools/health never hit the network.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
# Override DB dependency
|
||||
app.dependency_overrides[api_module.get_db] = _override_get_db
|
||||
|
||||
# Stub MCP client in ALL modules that imported it via `from mcp_client import`
|
||||
# (api + agent both have their own references that monkeypatch on the source
|
||||
# module alone won't affect).
|
||||
factory = lambda: stub_mcp
|
||||
monkeypatch.setattr(mcp_client_module, "get_mcp_client", factory)
|
||||
monkeypatch.setattr(api_module, "get_mcp_client", factory)
|
||||
monkeypatch.setattr(agent_module, "get_mcp_client", factory)
|
||||
monkeypatch.setattr(mcp_client_module, "close_mcp_client", AsyncMock_dummy())
|
||||
|
||||
transport = httpx.ASGITransport(app=app)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as client:
|
||||
yield client
|
||||
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def app_client_mcp_down(clean_db, monkeypatch, stub_mcp_unreachable) -> AsyncIterator:
|
||||
"""Same as app_client but MCP client raises CancelledError on health()."""
|
||||
import httpx
|
||||
|
||||
app.dependency_overrides[api_module.get_db] = _override_get_db
|
||||
factory = lambda: stub_mcp_unreachable
|
||||
monkeypatch.setattr(mcp_client_module, "get_mcp_client", factory)
|
||||
monkeypatch.setattr(api_module, "get_mcp_client", factory)
|
||||
monkeypatch.setattr(agent_module, "get_mcp_client", factory)
|
||||
monkeypatch.setattr(mcp_client_module, "close_mcp_client", AsyncMock_dummy())
|
||||
|
||||
transport = httpx.ASGITransport(app=app)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as client:
|
||||
yield client
|
||||
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
# === Seeding ===
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def seed_agent(clean_db) -> dict:
|
||||
"""Insert one db-source agent 'test_agent' for tests that need it."""
|
||||
payload = {
|
||||
"id": "test_agent",
|
||||
"name": "Test Agent",
|
||||
"description": "Seeded test agent",
|
||||
"system_prompt": "You are a test agent. Reply concisely.",
|
||||
"allowed_tools": ["echo", "calculate"],
|
||||
"model": None,
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 2000,
|
||||
"enabled": True,
|
||||
}
|
||||
await db_module.upsert_agent(clean_db, payload, source="db")
|
||||
return payload
|
||||
Reference in New Issue
Block a user