diff --git a/src/supervaizer/__init__.py b/src/supervaizer/__init__.py index 921614f..a293f09 100644 --- a/src/supervaizer/__init__.py +++ b/src/supervaizer/__init__.py @@ -78,6 +78,12 @@ "supervaizer.contracts", "SUPERVAIZER_V2_CONTRACT_VERSION", ), + "AGENT_CUSTOM_ACTION_PREFIX": ( + "supervaizer.contracts", + "AGENT_CUSTOM_ACTION_PREFIX", + ), + "AGENT_REFRESH_ACTION": ("supervaizer.contracts", "AGENT_REFRESH_ACTION"), + "AGENT_REFRESH_EFFECT": ("supervaizer.contracts", "AGENT_REFRESH_EFFECT"), "WORKSPACE_BINDING_CREATE_ACTION": ( "supervaizer.contracts", "WORKSPACE_BINDING_CREATE_ACTION", @@ -123,6 +129,8 @@ "V2A2UISubmitDefinition": ("supervaizer.contracts", "V2A2UISubmitDefinition"), "V2AgentCapabilities": ("supervaizer.contracts", "V2AgentCapabilities"), "V2AgentIdentity": ("supervaizer.contracts", "V2AgentIdentity"), + "V2AgentMethod": ("supervaizer.contracts", "V2AgentMethod"), + "V2AgentMethods": ("supervaizer.contracts", "V2AgentMethods"), "V2ArtifactRef": ("supervaizer.contracts", "V2ArtifactRef"), "V2ArtifactTypeDefinition": ( "supervaizer.contracts", diff --git a/src/supervaizer/agent.py b/src/supervaizer/agent.py index 1aa0ca7..e2f04be 100644 --- a/src/supervaizer/agent.py +++ b/src/supervaizer/agent.py @@ -14,6 +14,7 @@ import json import re from enum import Enum +from importlib import import_module from typing import ( TYPE_CHECKING, Any, @@ -36,7 +37,12 @@ from supervaizer.__version__ import VERSION from supervaizer.case import CaseNodes from supervaizer.common import ApiSuccess, SvBaseModel, log -from supervaizer.contracts import SupervaizerV2AgentRegistrationContract +from supervaizer.contracts import ( + SupervaizerV2AgentRegistrationContract, + V2ActionRequest, + V2AgentMethod, + V2AgentMethods, +) from supervaizer.data_resource import DataResource from supervaizer.event import JobStartConfirmationEvent from supervaizer.job import Job, JobContext, JobResponse @@ -695,6 +701,11 @@ class AgentAbstract(SvBaseModel): default=None, description="Optional Supervaizer v2 registration contract for A2A/A2UI Studio integrations", ) + v2_methods: V2AgentMethods | None = Field( + default=None, + description="Optional agent-level Supervaizer v2 method declarations", + exclude=True, + ) model_config = cast( ConfigDict, {"reference_group": "Core", "arbitrary_types_allowed": True} @@ -727,6 +738,7 @@ def __init__( supervaizer_v2_registration: SupervaizerV2AgentRegistrationContract | dict[str, Any] | None = None, + v2_methods: V2AgentMethods | dict[str, Any] | None = None, **kwargs: Any, ) -> None: """ @@ -792,10 +804,12 @@ def __init__( custom_routes=custom_routes, data_resources=data_resources or [], supervaizer_v2_registration=supervaizer_v2_registration, + v2_methods=v2_methods, **kwargs, ) self._validate_supervaizer_v2_identity() + self._apply_v2_method_capabilities() seen_resource_names: set[str] = set() for r in self.data_resources: @@ -821,6 +835,17 @@ def _validate_supervaizer_v2_identity(self) -> None: f"{declared_slug!r} != {self.slug!r}" ) + def _apply_v2_method_capabilities(self) -> None: + if self.supervaizer_v2_registration is None or self.v2_methods is None: + return + actions = [ + *self.supervaizer_v2_registration.capabilities.actions, + *self.v2_methods.action_ids, + ] + self.supervaizer_v2_registration.capabilities.actions = list( + dict.fromkeys(actions) + ) + @property def slug(self) -> str: return slugify(self.name) @@ -975,6 +1000,37 @@ def _declared_method_paths(self) -> set[str]: methods.extend(self.methods.custom.values()) return {method.method for method in methods if method is not None} + @property + def v2_action_ids(self) -> list[str]: + if self.v2_methods is None: + return [] + return self.v2_methods.action_ids + + def v2_method_for_action(self, action: str) -> V2AgentMethod | None: + if self.v2_methods is None: + return None + return self.v2_methods.method_for_action(action) + + def _declared_v2_method_paths(self) -> set[str]: + if self.v2_methods is None: + return set() + methods = [self.v2_methods.refresh, *self.v2_methods.custom.values()] + return {method.method for method in methods if method is not None} + + def execute_v2_action_method(self, action: str, request: V2ActionRequest) -> Any: + agent_method = self.v2_method_for_action(action) + if agent_method is None: + raise ValueError(f"Agent v2 action is not declared on agent: {action}") + if agent_method.method not in self._declared_v2_method_paths(): + raise ValueError( + f"Agent v2 method path is not declared on agent: {agent_method.method}" + ) + + module_name, func_name = agent_method.method.rsplit(".", 1) + module = import_module(module_name) + action_method = getattr(module, func_name) + return action_method(request=request, **agent_method.params) + def job_start( self, job: Job, diff --git a/src/supervaizer/contracts.py b/src/supervaizer/contracts.py index e0a6290..8c821a5 100644 --- a/src/supervaizer/contracts.py +++ b/src/supervaizer/contracts.py @@ -13,11 +13,12 @@ from __future__ import annotations +import re from collections.abc import Iterable from enum import StrEnum from typing import Any, Literal -from pydantic import BaseModel, Field, model_validator +from pydantic import BaseModel, Field, field_validator, model_validator CONTROLLER_CONTRACT_VERSION = "1.0" API_VERSION = "v1" @@ -28,6 +29,10 @@ WORKSPACE_BINDING_OPTIONS_ACTION = "workspace_binding.options" WORKSPACE_BINDING_CREATE_ACTION = "workspace_binding.create" WORKSPACE_BINDING_CREATE_SURFACE = "workspace_binding.create" +AGENT_REFRESH_ACTION = "agent.refresh" +AGENT_REFRESH_EFFECT = "agent.refreshed" +AGENT_CUSTOM_ACTION_PREFIX = "agent.custom." +_AGENT_CUSTOM_METHOD_KEY_RE = re.compile(r"^[A-Za-z0-9_-]+$") class ContractModel(BaseModel): @@ -310,6 +315,47 @@ class V2AgentCapabilities(ContractModel): artifact_types: list[V2ArtifactTypeDefinition] = Field(default_factory=list) +class V2AgentMethod(ContractModel): + method: str + params: dict[str, Any] = Field(default_factory=dict) + description: str | None = None + is_async: bool = False + timeout: int | None = 600 + + +class V2AgentMethods(ContractModel): + refresh: V2AgentMethod | None = None + custom: dict[str, V2AgentMethod] = Field(default_factory=dict) + + @field_validator("custom") + @classmethod + def validate_custom_method_names( + cls, value: dict[str, V2AgentMethod] + ) -> dict[str, V2AgentMethod]: + for name in value: + if not _AGENT_CUSTOM_METHOD_KEY_RE.fullmatch(name): + raise ValueError( + "agent custom method keys may only contain letters, numbers, " + "underscores, and hyphens" + ) + return value + + @property + def action_ids(self) -> list[str]: + actions: list[str] = [] + if self.refresh is not None: + actions.append(AGENT_REFRESH_ACTION) + actions.extend(f"{AGENT_CUSTOM_ACTION_PREFIX}{name}" for name in self.custom) + return actions + + def method_for_action(self, action: str) -> V2AgentMethod | None: + if action == AGENT_REFRESH_ACTION: + return self.refresh + if action.startswith(AGENT_CUSTOM_ACTION_PREFIX): + return self.custom.get(action.removeprefix(AGENT_CUSTOM_ACTION_PREFIX)) + return None + + class V2JobSyncPolicy(ContractModel): action: str = "job.sync" supported_statuses: list[str] = Field(default_factory=list) @@ -524,6 +570,7 @@ def build_v2_agent_registration( datasets: Iterable[V2DatasetDefinition | dict[str, Any]] = (), dashboards: Iterable[V2DashboardDefinition | dict[str, Any]] = (), workspace_binding: V2WorkspaceBindingDefinition | dict[str, Any] | None = None, + agent_methods: V2AgentMethods | dict[str, Any] | None = None, case_lanes: Iterable[V2CaseLaneDefinition | dict[str, Any]] = (), artifact_types: Iterable[V2ArtifactTypeDefinition | dict[str, Any]] = (), job_policy: V2JobPolicy | dict[str, Any] | None = None, @@ -538,6 +585,7 @@ def build_v2_agent_registration( dataset_definitions = _contract_list(datasets, V2DatasetDefinition) dashboard_definitions = _contract_list(dashboards, V2DashboardDefinition) workspace_binding_definition = _workspace_binding(workspace_binding) + agent_method_definitions = _agent_methods(agent_methods) sync_policy = _job_policy(job_policy) capability_surfaces = _unique_strings([ @@ -553,6 +601,7 @@ def build_v2_agent_registration( *_dataset_action_ids(dataset_definitions), *(_job_sync_actions(sync_policy)), *_workspace_binding_action_ids(workspace_binding_definition), + *_agent_method_action_ids(agent_method_definitions), ]) return SupervaizerV2AgentRegistrationContract( @@ -623,6 +672,16 @@ def _workspace_binding( return V2WorkspaceBindingDefinition.model_validate(value) +def _agent_methods( + value: V2AgentMethods | dict[str, Any] | None, +) -> V2AgentMethods | None: + if value is None: + return None + if isinstance(value, V2AgentMethods): + return value + return V2AgentMethods.model_validate(value) + + def _unique_strings(values: Iterable[str]) -> list[str]: seen: set[str] = set() result: list[str] = [] @@ -672,6 +731,12 @@ def _job_sync_actions(job_policy: V2JobPolicy) -> list[str]: return [job_policy.sync.action] +def _agent_method_action_ids(agent_methods: V2AgentMethods | None) -> list[str]: + if agent_methods is None: + return [] + return agent_methods.action_ids + + def _workspace_binding_action_ids( workspace_binding: V2WorkspaceBindingDefinition | None, ) -> list[str]: diff --git a/src/supervaizer/server.py b/src/supervaizer/server.py index 1ff6137..db1393d 100644 --- a/src/supervaizer/server.py +++ b/src/supervaizer/server.py @@ -98,6 +98,13 @@ SCHEDULED_STEP_SHUTDOWN_TIMEOUT_SECONDS = 5.0 +def _agent_v2_method_handler(agent: Agent, action: str) -> ActionHandler: + def handler(request: Any) -> Any: + return agent.execute_v2_action_method(action, request) + + return handler + + class ServerAbstract(SvBaseModel): """ API Server for the Supervaize Controller. @@ -445,6 +452,7 @@ async def validation_exception_handler( # Store server instance on app state before building routers self.app.state.server = self # <-- MOVED earlier (was after route mount) + self._register_agent_v2_method_handlers() # Activate API + A2A routes when supervisor account or local mode is set if self.supervisor_account or local_mode: @@ -659,6 +667,15 @@ def register_v2_action( register_v2_action_handler(self, action, handler, agent_slug=agent_slug) return handler + def _register_agent_v2_method_handlers(self) -> None: + for agent in self.agents: + for action in agent.v2_action_ids: + self.register_v2_action( + action, + _agent_v2_method_handler(agent, action), + agent_slug=agent.slug, + ) + def v2_action( self, action: str, *, agent_slug: str | None = None ) -> Callable[[ActionHandler], ActionHandler]: diff --git a/tests/fixtures/supervaizer_v2/agent_interviewer_mvp.json b/tests/fixtures/supervaizer_v2/agent_interviewer_mvp.json index 3d2ccda..b923eb8 100644 --- a/tests/fixtures/supervaizer_v2/agent_interviewer_mvp.json +++ b/tests/fixtures/supervaizer_v2/agent_interviewer_mvp.json @@ -46,7 +46,7 @@ "job.start", "job.stop", "job.sync", - "campaigns.sync", + "agent.refresh", "artifact.get", "resource.campaigns.list", "resource.campaigns.create", diff --git a/tests/test_a2a.py b/tests/test_a2a.py index 8f704b3..1b00411 100644 --- a/tests/test_a2a.py +++ b/tests/test_a2a.py @@ -21,7 +21,15 @@ from cryptography.hazmat.primitives.asymmetric import ed25519, rsa from fastapi.testclient import TestClient -from supervaizer import Agent, Server +from supervaizer import ( + AGENT_REFRESH_ACTION, + AGENT_REFRESH_EFFECT, + Agent, + Server, + V2AgentMethod, + V2AgentMethods, + build_v2_agent_registration, +) from supervaizer.access import API_KEYS from supervaizer.contracts import ( V2ActionRequest, @@ -56,6 +64,13 @@ from supervaizer.workspace_authorization import WORKSPACE_AUTHORIZATION_HEADER +def _test_agent_refresh(request: V2ActionRequest) -> dict[str, object]: + return { + "status": "ok", + "effects": [{"type": AGENT_REFRESH_EFFECT, "request_id": request.request_id}], + } + + def _a2a_write_headers(server: Server) -> dict[str, str]: return {"X-API-Key": server.api_key or ""} @@ -1622,6 +1637,57 @@ def preview_job_start(request: V2ActionRequest) -> dict[str, object]: } +def test_server_registers_agent_v2_method_handlers() -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + agent = Agent( + name="Agent Name", + version="1.0.0", + supervaizer_v2_registration=build_v2_agent_registration( + agent_id="agent-name", + agent_slug="agent-name", + display_name="Agent Name", + agent_card_url="/.well-known/agents/v1.0.0/agent-name_agent.json", + controller_url="/a2a", + a2ui_catalog_version="test.0", + ), + v2_methods=V2AgentMethods( + refresh=V2AgentMethod(method="tests.test_a2a._test_agent_refresh") + ), + ) + server = Server( + agents=[agent], + private_key=private_key, + api_key="test-api-key", + admin_interface=False, + ) + headers = _authorized_a2a_headers( + server, + agent_slug=agent.slug, + scopes=[SUPERVAIZER_ACTION_INVOKE_METHOD, AGENT_REFRESH_ACTION], + ) + client = TestClient(server.app) + + response = client.post( + "/a2a", + headers=headers, + json={ + "jsonrpc": "2.0", + "id": "rpc-agent-refresh", + "method": SUPERVAIZER_ACTION_INVOKE_METHOD, + "params": _v2_action_payload( + action=AGENT_REFRESH_ACTION, + agent_slug=agent.slug, + ), + }, + ) + + assert response.status_code == 200 + assert response.json()["result"] == { + "status": "ok", + "effects": [{"type": AGENT_REFRESH_EFFECT, "request_id": "request-1"}], + } + + def test_server_v2_surface_decorator_registers_handler(server_fixture: Server) -> None: agent_slug = server_fixture.agents[0].slug headers = _authorized_a2a_headers( diff --git a/tests/test_agent.py b/tests/test_agent.py index f27c41e..318752e 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -19,7 +19,16 @@ import pytest from pydantic import BaseModel, ValidationError -from supervaizer import Agent, AgentMethod, AgentMethods, ApiSuccess, Server +from supervaizer import ( + AGENT_REFRESH_ACTION, + Agent, + AgentMethod, + AgentMethods, + ApiSuccess, + Server, + V2AgentMethod, + V2AgentMethods, +) from supervaizer.agent import AgentMethodField, AgentMethodsAbstract, FieldTypeEnum from supervaizer.job import Job, JobContext, JobResponse from supervaizer.lifecycle import EntityStatus @@ -1237,6 +1246,39 @@ def test_agent_accepts_v2_registration_with_matching_slug() -> None: assert agent.supervaizer_v2_registration.agent.slug == "agent-name" +def test_agent_v2_methods_are_added_to_v2_capabilities() -> None: + agent = Agent( + name="Agent Name", + version="1.0.0", + supervaizer_v2_registration={ + "agent": { + "id": "agent_name", + "slug": "agent-name", + "display_name": "Agent Name", + }, + "versions": { + "a2ui_version": "v0.8", + "a2ui_catalog_version": "test.0", + "a2a_version": "0.2.6", + }, + "a2a": { + "agent_card_url": "/.well-known/agents/v1.0.0/agent-name_agent.json", + "controller_url": "/a2a", + }, + "capabilities": {"actions": ["job.sync"]}, + }, + v2_methods=V2AgentMethods( + refresh=V2AgentMethod(method="tests.test_agent.refresh_agent") + ), + ) + + assert agent.supervaizer_v2_registration is not None + assert agent.supervaizer_v2_registration.capabilities.actions == [ + "job.sync", + AGENT_REFRESH_ACTION, + ] + + def test_agent_rejects_v2_registration_with_mismatched_slug() -> None: """Agent rejects a v2 registration whose declared slug differs from runtime slug.""" with pytest.raises( diff --git a/tests/test_contracts.py b/tests/test_contracts.py index 79f38ee..5cbe2d2 100644 --- a/tests/test_contracts.py +++ b/tests/test_contracts.py @@ -17,6 +17,7 @@ from pydantic import ValidationError from supervaizer.contracts import ( + AGENT_REFRESH_ACTION, API_VERSION, AgentMethodsContract, ControllerEndpoint, @@ -27,6 +28,8 @@ V2A2UIResourceImportDocument, V2ActionRequest, V2ActionResult, + V2AgentMethod, + V2AgentMethods, V2AwaitingState, V2CaseSnapshot, V2DashboardWidgetDataRef, @@ -262,7 +265,7 @@ def test_v2_agent_interviewer_registration_fixture() -> None: assert ( "mission.agent.surface.scenario_builder" in registration.capabilities.surfaces ) - assert "campaigns.sync" in registration.capabilities.actions + assert AGENT_REFRESH_ACTION in registration.capabilities.actions assert "resource.campaign_contacts.create" in registration.capabilities.actions assert "resource.campaign_contacts.delete" in registration.capabilities.actions assert "resource.contacts.import" in registration.capabilities.actions @@ -454,6 +457,30 @@ def test_build_v2_agent_registration_derives_capabilities() -> None: assert registration.capabilities.case_lanes[0].default is True +def test_build_v2_agent_registration_derives_agent_method_actions() -> None: + registration = build_v2_agent_registration( + agent_id="hello", + agent_slug="hello-world", + display_name="Hello World", + agent_card_url="/.well-known/agents/v1/hello-world_agent.json", + controller_url="/a2a", + a2ui_catalog_version="supervaizer-v2-local.0", + agent_methods=V2AgentMethods( + refresh=V2AgentMethod(method="hello_agent.refresh"), + custom={ + "reindex": V2AgentMethod(method="hello_agent.reindex"), + "dry-run": V2AgentMethod(method="hello_agent.dry_run"), + }, + ), + ) + + assert registration.capabilities.actions == [ + AGENT_REFRESH_ACTION, + "agent.custom.reindex", + "agent.custom.dry-run", + ] + + def test_v2_workspace_binding_required_requires_mode() -> None: with pytest.raises(ValidationError, match="at least one mode"): V2WorkspaceBindingDefinition(required=True) @@ -671,6 +698,8 @@ def test_v2_contract_models_are_public_sdk_exports() -> None: assert supervaizer.V2A2ATransport.__name__ == "V2A2ATransport" assert supervaizer.V2A2UIResourceImportDocument is V2A2UIResourceImportDocument assert supervaizer.V2ActionRequest is V2ActionRequest + assert supervaizer.V2AgentMethod is V2AgentMethod + assert supervaizer.V2AgentMethods is V2AgentMethods assert supervaizer.V2DashboardDefinition.__name__ == "V2DashboardDefinition" assert supervaizer.V2DashboardWidgetDataRef is V2DashboardWidgetDataRef assert supervaizer.V2DashboardWidgetDefinition is V2DashboardWidgetDefinition