83 lines
3.5 KiB
Python
83 lines
3.5 KiB
Python
from unittest.mock import AsyncMock, Mock
|
|||
|
|
|
||
|
|
import pytest
|
||
|
|
|
||
|
|
from app.core.contracts import AgentDefinition, AgentRequest, CoreResult, RequestContext
|
||
|
|
from app.core.errors import ForbiddenAgentError, ValidationAgentError
|
||
|
|
from app.service.agent.base import BaseAgent
|
||
|
|
from app.service.agent.factory import AgentFactory
|
||
|
|
from app.service.agent_run_application_service import AgentRunApplicationService
|
||
|
|
|
||
|
|
|
||
|
|
class RiskAgent(BaseAgent):
|
||
|
|
async def handle(self, request: AgentRequest, context: RequestContext) -> CoreResult:
|
||
|
|
return CoreResult(text="safe")
|
||
|
|
|
||
|
|
|
||
|
|
def registry() -> tuple[AgentFactory, Mock]:
|
||
|
|
definition = AgentDefinition(agent_type="risk", version="1",
|
||
|
|
allowed_roles=("risk_operator",), allowed_portals=("api",))
|
||
|
|
builder = Mock(side_effect=lambda _: RiskAgent(definition))
|
||
|
|
factory = AgentFactory()
|
||
|
|
factory.register(definition, builder)
|
||
|
|
return factory, builder
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.parametrize("roles,portal,permissions", [
|
||
|
|
(("customer",), "api", ("agent:run",)),
|
||
|
|
(("risk_operator",), "forged", ("agent:run",)),
|
||
|
|
(("risk_operator",), "api", ()),
|
||
|
|
((), "api", ("agent:run",)),
|
||
|
|
])
|
||
|
|
def test_denied_before_builder(roles, portal, permissions):
|
||
|
|
factory, builder = registry()
|
||
|
|
with pytest.raises(ForbiddenAgentError):
|
||
|
|
factory.create("risk", RequestContext(user_id="1", trace_id="t",
|
||
|
|
roles=roles, portal=portal, permissions=permissions))
|
||
|
|
builder.assert_not_called()
|
||
|
|
|
||
|
|
|
||
|
|
def test_authorized_instances_are_not_shared():
|
||
|
|
factory, _ = registry()
|
||
|
|
context = RequestContext(user_id="1", trace_id="t", roles=("risk_operator",),
|
||
|
|
permissions=("agent:run",))
|
||
|
|
assert factory.create("risk", context) is not factory.create("risk", context)
|
||
|
|
|
||
|
|
|
||
|
|
async def test_direct_execution_still_checks_authorization():
|
||
|
|
definition = AgentDefinition(agent_type="risk", version="1",
|
||
|
|
allowed_roles=("risk_operator",), allowed_portals=("api",))
|
||
|
|
request = AgentRequest(agent_type="risk", message="test", session_id="s",
|
||
|
|
idempotency_key="1234567890123456")
|
||
|
|
with pytest.raises(ForbiddenAgentError):
|
||
|
|
_ = [e async for e in RiskAgent(definition).execute(
|
||
|
|
request, RequestContext(user_id="1", trace_id="t"), "r")]
|
||
|
|
|
||
|
|
|
||
|
|
async def test_accept_denial_only_writes_audit():
|
||
|
|
factory, builder = registry()
|
||
|
|
session = AsyncMock()
|
||
|
|
session.begin = Mock(return_value=AsyncMock())
|
||
|
|
session.add = Mock()
|
||
|
|
request = AgentRequest(agent_type="risk", message="test", session_id="s",
|
||
|
|
idempotency_key="1234567890123456")
|
||
|
|
with pytest.raises(ForbiddenAgentError):
|
||
|
|
await AgentRunApplicationService(session, factory).accept(
|
||
|
|
request, RequestContext(user_id="1", trace_id="t", roles=("customer",)))
|
||
|
|
assert session.add.call_count == 1
|
||
|
|
assert session.add.call_args.args[0].action_type == "agent.access_denied"
|
||
|
|
session.flush.assert_not_awaited()
|
||
|
|
builder.assert_not_called()
|
||
|
|
|
||
|
|
|
||
|
|
def test_builder_cannot_swap_authorized_definition():
|
||
|
|
factory = AgentFactory()
|
||
|
|
definition = AgentDefinition(agent_type="risk", version="1",
|
||
|
|
allowed_roles=("risk_operator",), allowed_portals=("api",))
|
||
|
|
factory.register(
|
||
|
|
definition, lambda _: RiskAgent(AgentDefinition(agent_type="other", version="1"))
|
||
|
|
)
|
||
|
|
with pytest.raises(ValidationAgentError):
|
||
|
|
factory.create("risk", RequestContext(user_id="1", trace_id="t",
|
||
|
|
roles=("risk_operator",), permissions=("agent:run",)))
|