diff --git a/docs/CHANGELOG.md b/docs/CHANGELOG.md index ba52c54..26e0ffe 100644 --- a/docs/CHANGELOG.md +++ b/docs/CHANGELOG.md @@ -17,6 +17,33 @@ All notable changes to this project will be documented in this file. ## [Unreleased] +### Added + +- **Controller contract metadata for Studio** — `server.register` now advertises `controller_contract_version`, `/api` base path, and canonical endpoint templates so Studio can route without hardcoded Supervaizer paths. The same contract is available at `GET /api/supervaizer/contract`. +- **DataResource request context** — DataResource callbacks can opt into a `context` keyword with workspace, mission, agent, and request identifiers while legacy callbacks continue to work unchanged. +- **CaseNodes update helpers** — `CaseNodes.node_index()` and `CaseNodes.make_update()` derive step indexes from the declared node set instead of requiring consumers to maintain a separate index map. + +### Changed + +- **DataResource fields can be marked sensitive** — `DataResourceField.sensitive` is included in registration payloads for Studio masking. + +### Tests + +- New: `tests/test_contracts.py` for contract endpoint templates, schema export, import safety, resolver behavior, and DataResource context headers. +- New: `tests/test_case.py` coverage for `CaseNodes.node_index()` and `CaseNodes.make_update()`. +- Updated: `tests/test_routes.py` to verify DataResource callbacks receive `DataResourceContext` from forwarded `X-Supervaize-*` headers. +- Updated: `tests/test_server.py` to assert registration info and `/api/supervaizer/contract` expose matching `/api` contract metadata. +- Deleted: none. + +`just test` + +| Status | Count | +| ---------- | ----- | +| ✅ Passed | 514 | +| 🤔 Skipped | 0 | +| 🔴 Failed | 0 | +| ⏱️ in | 01:03 | + ## [0.15.1] - 2026-04-20 ### Fixed diff --git a/src/supervaizer/__init__.py b/src/supervaizer/__init__.py index 7c1cbf4..32a0677 100644 --- a/src/supervaizer/__init__.py +++ b/src/supervaizer/__init__.py @@ -4,113 +4,114 @@ # If a copy of the MPL was not distributed with this file, you can obtain one at # https://mozilla.org/MPL/2.0/. -# Copyright (c) 2024-2025 Alain Prasquier - Supervaize.com. All rights reserved. -# -# This Source Code Form is subject to the terms of the Mozilla Public License, v. 2.0. -# If a copy of the MPL was not distributed with this file, you can obtain one at -# https://mozilla.org/MPL/2.0/. +"""Public Supervaizer SDK surface. + +Exports are resolved lazily so contract-only consumers such as Studio can import +``supervaizer.contracts`` without loading the controller runtime. +""" + +from __future__ import annotations + +from importlib import import_module +from typing import Any + +_EXPORTS: dict[str, tuple[str, str | None]] = { + "protocol": ("supervaizer.protocol", None), + "Account": ("supervaizer.account", "Account"), + "Agent": ("supervaizer.agent", "Agent"), + "AgentCustomMethodParams": ("supervaizer.agent", "AgentCustomMethodParams"), + "AgentMethod": ("supervaizer.agent", "AgentMethod"), + "AgentMethodField": ("supervaizer.agent", "AgentMethodField"), + "AgentMethodParams": ("supervaizer.agent", "AgentMethodParams"), + "AgentMethods": ("supervaizer.agent", "AgentMethods"), + "AgentResponse": ("supervaizer.agent", "AgentResponse"), + "FieldTypeEnum": ("supervaizer.agent", "FieldTypeEnum"), + "ApiError": ("supervaizer.common", "ApiError"), + "ApiResult": ("supervaizer.common", "ApiResult"), + "ApiSuccess": ("supervaizer.common", "ApiSuccess"), + "Case": ("supervaizer.case", "Case"), + "CaseNode": ("supervaizer.case", "CaseNode"), + "CaseNodes": ("supervaizer.case", "CaseNodes"), + "CaseNodeType": ("supervaizer.case", "CaseNodeType"), + "CaseNodeUpdate": ("supervaizer.case", "CaseNodeUpdate"), + "Cases": ("supervaizer.case", "Cases"), + "DataResource": ("supervaizer.data_resource", "DataResource"), + "DataResourceContext": ("supervaizer.data_resource", "DataResourceContext"), + "DataResourceField": ("supervaizer.data_resource", "DataResourceField"), + "Editable": ("supervaizer.data_resource", "Editable"), + "FieldType": ("supervaizer.data_resource", "FieldType"), + "AgentRegisterEvent": ("supervaizer.event", "AgentRegisterEvent"), + "CaseStartEvent": ("supervaizer.event", "CaseStartEvent"), + "CaseUpdateEvent": ("supervaizer.event", "CaseUpdateEvent"), + "Event": ("supervaizer.event", "Event"), + "JobFinishedEvent": ("supervaizer.event", "JobFinishedEvent"), + "JobStartConfirmationEvent": ("supervaizer.event", "JobStartConfirmationEvent"), + "ServerRegisterEvent": ("supervaizer.event", "ServerRegisterEvent"), + "EntityEvents": ("supervaizer.lifecycle", "EntityEvents"), + "EntityLifecycle": ("supervaizer.lifecycle", "EntityLifecycle"), + "EntityStatus": ("supervaizer.lifecycle", "EntityStatus"), + "Job": ("supervaizer.job", "Job"), + "JobContext": ("supervaizer.job", "JobContext"), + "JobInstructions": ("supervaizer.job", "JobInstructions"), + "JobResponse": ("supervaizer.job", "JobResponse"), + "Jobs": ("supervaizer.job", "Jobs"), + "Parameter": ("supervaizer.parameter", "Parameter"), + "ParametersSetup": ("supervaizer.parameter", "ParametersSetup"), + "Server": ("supervaizer.server", "Server"), + "create_error_response": ("supervaizer.server_utils", "create_error_response"), + "ErrorResponse": ("supervaizer.server_utils", "ErrorResponse"), + "ErrorType": ("supervaizer.server_utils", "ErrorType"), + "Telemetry": ("supervaizer.telemetry", "Telemetry"), + "TelemetryCategory": ("supervaizer.telemetry", "TelemetryCategory"), + "TelemetrySeverity": ("supervaizer.telemetry", "TelemetrySeverity"), + "TelemetryType": ("supervaizer.telemetry", "TelemetryType"), + "AgentMethodContract": ("supervaizer.contracts", "AgentMethodContract"), + "AgentMethodsContract": ("supervaizer.contracts", "AgentMethodsContract"), + "AgentRegistrationContract": ("supervaizer.contracts", "AgentRegistrationContract"), + "ControllerContract": ("supervaizer.contracts", "ControllerContract"), + "ControllerEndpoint": ("supervaizer.contracts", "ControllerEndpoint"), + "DataResourceContract": ("supervaizer.contracts", "DataResourceContract"), + "DataResourceFieldContract": ("supervaizer.contracts", "DataResourceFieldContract"), + "DynamicChoicesRequest": ("supervaizer.contracts", "DynamicChoicesRequest"), + "DynamicChoicesResponse": ("supervaizer.contracts", "DynamicChoicesResponse"), + "EventType": ("supervaizer.contracts", "EventType"), + "JobStartRequest": ("supervaizer.contracts", "JobStartRequest"), + "ServerRegistrationContract": ( + "supervaizer.contracts", + "ServerRegistrationContract", + ), + "build_data_resource_context_headers": ( + "supervaizer.contracts", + "build_data_resource_context_headers", + ), + "controller_contract_info": ("supervaizer.contracts", "controller_contract_info"), + "resolve_controller_endpoint": ( + "supervaizer.contracts", + "resolve_controller_endpoint", + ), +} + +__all__ = sorted(_EXPORTS) -from supervaizer import protocol -from supervaizer.account import Account -from supervaizer.agent import ( - Agent, - AgentCustomMethodParams, - AgentMethod, - AgentMethodParams, - AgentMethods, - AgentMethodField, - AgentResponse, - FieldTypeEnum, -) -from supervaizer.case import ( - Case, - CaseNodeUpdate, - CaseNodeType, - Cases, - CaseNode, - CaseNodes, -) -from supervaizer.common import ApiError, ApiResult, ApiSuccess -from supervaizer.data_resource import ( - DataResource, - DataResourceField, - Editable, - FieldType, -) -from supervaizer.event import ( - AgentRegisterEvent, - CaseStartEvent, - CaseUpdateEvent, - Event, - EventType, - JobFinishedEvent, - JobStartConfirmationEvent, - ServerRegisterEvent, -) -from supervaizer.job import Job, JobContext, JobInstructions, JobResponse, Jobs -from supervaizer.lifecycle import EntityEvents, EntityLifecycle, EntityStatus -from supervaizer.parameter import Parameter, ParametersSetup -from supervaizer.server import Server -from supervaizer.server_utils import ErrorResponse, ErrorType, create_error_response -from supervaizer.telemetry import ( - Telemetry, - TelemetryCategory, - TelemetrySeverity, - TelemetryType, -) +def _rebuild_forward_refs() -> None: + """Resolve forward refs for models that depend on the public import order.""" + try: + case_module = import_module("supervaizer.case") + agent_module = import_module("supervaizer.agent") + case_module.Case.model_rebuild() + agent_module.AgentResponse.model_rebuild() + except Exception: + pass -__all__ = [ - "Account", - "Agent", - "AgentCustomMethodParams", - "AgentMethod", - "AgentMethodField", - "AgentMethodParams", - "AgentMethods", - "AgentRegisterEvent", - "ApiError", - "ApiResult", - "ApiSuccess", - "Case", - "CaseNode", - "CaseNodes", - "CaseNodeType", - "CaseNodeUpdate", - "Cases", - "CaseStartEvent", - "CaseUpdateEvent", - "create_error_response", - "DataResource", - "DataResourceField", - "Editable", - "FieldType", - "EntityEvents", - "EntityLifecycle", - "EntityStatus", - "ErrorResponse", - "ErrorType", - "Event", - "EventType", - "FieldTypeEnum", - "Job", - "JobContext", - "JobFinishedEvent", - "JobInstructions", - "JobResponse", - "Jobs", - "JobStartConfirmationEvent", - "Parameter", - "ParametersSetup", - "protocol", - "Server", - "ServerRegisterEvent", - "Telemetry", - "TelemetryCategory", - "TelemetrySeverity", - "TelemetryType", -] -# Rebuild models to resolve forward references after all imports are done -Case.model_rebuild() -AgentResponse.model_rebuild() +def __getattr__(name: str) -> Any: + if name not in _EXPORTS: + raise AttributeError(f"module 'supervaizer' has no attribute {name!r}") + module_name, attr_name = _EXPORTS[name] + module = import_module(module_name) + value = module if attr_name is None else getattr(module, attr_name) + globals()[name] = value + if module_name in {"supervaizer.agent", "supervaizer.case"}: + _rebuild_forward_refs() + return value diff --git a/src/supervaizer/agent.py b/src/supervaizer/agent.py index 1fd8bc7..f0226c7 100644 --- a/src/supervaizer/agent.py +++ b/src/supervaizer/agent.py @@ -931,11 +931,11 @@ def job_start( if job_response.payload: job.metadata = job_response.payload # Send confirmation event now that metadata is populated - event = JobStartConfirmationEvent( - job=job, - account=server.supervisor_account, - ) if server.supervisor_account is not None: + event = JobStartConfirmationEvent( + job=job, + account=server.supervisor_account, + ) server.supervisor_account.send_event(sender=job, event=event) else: log.warning( diff --git a/src/supervaizer/case.py b/src/supervaizer/case.py index 72ad479..d26e781 100644 --- a/src/supervaizer/case.py +++ b/src/supervaizer/case.py @@ -201,6 +201,33 @@ class CaseNodes(SvBaseModel): def get(self, name: str) -> CaseNode | None: return next((node for node in self.nodes if node.name == name), None) + def node_index(self, name: str, *, start: int = 1) -> int: + """Return the stable 1-based index for a named case node.""" + for offset, node in enumerate(self.nodes): + if node.name == name: + return start + offset + raise ValueError(f"Case node {name!r} not found") + + def make_update( + self, + name: str, + *, + payload: Dict[str, Any] | None = None, + cost: float = 0.0, + is_final: bool = False, + upsert_to: str | None = None, + ) -> CaseNodeUpdate: + """Build a CaseNodeUpdate with an index derived from this node set.""" + target_name = upsert_to or name + return CaseNodeUpdate( + name=name, + cost=cost, + payload=payload or {}, + is_final=is_final, + index=self.node_index(target_name), + upsert=upsert_to is not None, + ) + @property def registration_info(self) -> Dict[str, Any]: """Returns registration info for the case nodes""" @@ -520,3 +547,13 @@ def get_due_scheduled_steps(self) -> list[tuple]: def __contains__(self, case_id: str) -> bool: """Check if case exists in any job's registry""" return any(case_id in cases for cases in self.cases_by_job.values()) + + +def rebuild_case_model_forward_refs() -> None: + """Resolve runtime-only forward references used by Case models.""" + from supervaizer.account import Account + + Case.model_rebuild(_types_namespace={"Account": Account}) + + +rebuild_case_model_forward_refs() diff --git a/src/supervaizer/contracts.py b/src/supervaizer/contracts.py new file mode 100644 index 0000000..419be78 --- /dev/null +++ b/src/supervaizer/contracts.py @@ -0,0 +1,290 @@ +# Copyright (c) 2024-2026 Alain Prasquier - Supervaize.com. All rights reserved. +# +# This Source Code Form is subject to the terms of the Mozilla Public License, v. 2.0. +# If a copy of the MPL was not distributed with this file, you can obtain one at +# https://mozilla.org/MPL/2.0/. + +"""Versioned controller contract shared with Studio integrations. + +This module is intentionally import-light. Studio imports it directly as the +single source of truth for the controller wire contract, so it must not import +the Supervaizer server/runtime surface. +""" + +from __future__ import annotations + +from enum import StrEnum +from typing import Any + +from pydantic import BaseModel, Field + +CONTROLLER_CONTRACT_VERSION = "1.0" +API_BASE_PATH = "/api" + + +class ContractModel(BaseModel): + """Base class for SDK-owned wire contract models.""" + + model_config = {"use_enum_values": True, "extra": "allow"} + + +class ControllerEndpoint(StrEnum): + POST_AGENT_JOB_START = "POST_AGENT_JOB_START" + POST_AGENT_JOB_CUSTOM = "POST_AGENT_JOB_CUSTOM" + GET_JOB_STATUS = "GET_JOB_STATUS" + GET_AGENT_JOB_STATUS = "GET_AGENT_JOB_STATUS" + POST_AGENT_STOP = "POST_AGENT_STOP" + GET_AGENT_JOB_LIST = "GET_AGENT_JOB_LIST" + GET_AGENT_BY_ID = "GET_AGENT_BY_ID" + GET_AGENT_LIST = "GET_AGENT_LIST" + GET_AGENT_BY_SLUG = "GET_AGENT_BY_SLUG" + POST_AGENT_PARAMETERS = "POST_AGENT_PARAMETERS" + POST_AGENT_CASE_UPDATE = "POST_AGENT_CASE_UPDATE" + POST_AGENT_PARAMETER_VALIDATION = "POST_AGENT_PARAMETER_VALIDATION" + POST_AGENT_METHOD_FIELD_VALIDATION = "POST_AGENT_METHOD_FIELD_VALIDATION" + POST_AGENT_JOB_START_DYNAMIC_CHOICES = "POST_AGENT_JOB_START_DYNAMIC_CHOICES" + DATA_RESOURCE = "DATA_RESOURCE" + DATA_RESOURCE_ITEM = "DATA_RESOURCE_ITEM" + DATA_RESOURCE_IMPORT = "DATA_RESOURCE_IMPORT" + HEALTH_CHECK = "HEALTH_CHECK" + CONTROLLER_CONTRACT = "CONTROLLER_CONTRACT" + + +CONTROLLER_ENDPOINTS: dict[ControllerEndpoint, str] = { + ControllerEndpoint.POST_AGENT_JOB_START: "/api/supervaizer/agents/{agent_slug}/jobs", + ControllerEndpoint.POST_AGENT_JOB_CUSTOM: "/api/supervaizer/agents/{agent_slug}/custom/{method_name}", + ControllerEndpoint.GET_JOB_STATUS: "/api/supervaizer/jobs/{job_id}", + ControllerEndpoint.GET_AGENT_JOB_STATUS: "/api/supervaizer/agents/{agent_slug}/jobs/{job_id}", + ControllerEndpoint.POST_AGENT_STOP: "/api/supervaizer/agents/{agent_slug}/stop", + ControllerEndpoint.GET_AGENT_JOB_LIST: "/api/supervaizer/agents/{agent_slug}/jobs", + ControllerEndpoint.GET_AGENT_BY_ID: "/api/supervaizer/agents/{agent_id}", + ControllerEndpoint.GET_AGENT_LIST: "/api/supervaizer/agents", + ControllerEndpoint.GET_AGENT_BY_SLUG: "/api/supervaizer/agents/{agent_slug}", + ControllerEndpoint.POST_AGENT_PARAMETERS: "/api/supervaizer/agents/{agent_slug}/parameters", + ControllerEndpoint.POST_AGENT_CASE_UPDATE: "/api/supervaizer/jobs/{job_id}/cases/{case_id}/update", + ControllerEndpoint.POST_AGENT_PARAMETER_VALIDATION: "/api/supervaizer/agents/{agent_slug}/validate-agent-parameters", + ControllerEndpoint.POST_AGENT_METHOD_FIELD_VALIDATION: "/api/supervaizer/agents/{agent_slug}/validate-method-fields", + ControllerEndpoint.POST_AGENT_JOB_START_DYNAMIC_CHOICES: "/api/supervaizer/agents/{agent_slug}/start/dynamic_choices", + ControllerEndpoint.DATA_RESOURCE: "/api/agents/{agent_slug}/data/{resource_name}/", + ControllerEndpoint.DATA_RESOURCE_ITEM: "/api/agents/{agent_slug}/data/{resource_name}/{item_id}", + ControllerEndpoint.DATA_RESOURCE_IMPORT: "/api/agents/{agent_slug}/data/{resource_name}/import/", + ControllerEndpoint.HEALTH_CHECK: ".well-known/health", + ControllerEndpoint.CONTROLLER_CONTRACT: "/api/supervaizer/contract", +} + + +class EventType(StrEnum): + SERVER_REGISTER = "server.register" + SERVER_ONLINE = "server.online" + SERVER_DOWN = "server.down" + AGENT_REGISTER = "agent.register" + AGENT_WAKEUP = "agent.wakeup" + AGENT_ANOMALY = "agent.anomaly" + AGENT_SEND_ANOMALY = "agent.anomaly" + AGENT_PING = "agent.ping" + INTERMEDIARY = "agent.intermediary" + JOB_START = "agent.job.start" + JOB_START_CONFIRMATION = "agent.job.start.confirmation" + JOB_END = "agent.job.end" + JOB_STATUS = "agent.job.status" + JOB_RESULT = "agent.job.result" + JOB_ERROR = "agent.job.error" + JOB_TIMEOUT = "agent.job.timeout" + CASE_START = "agent.case.start" + CASE_END = "agent.case.end" + CASE_STATUS = "agent.case.status" + CASE_RESULT = "agent.case.result" + CASE_UPDATE = "agent.case.update" + CASE_ERROR = "agent.case.error" + + +class DataResourceContextContract(ContractModel): + workspace_id: str | None = None + workspace_slug: str | None = None + mission_id: str | None = None + agent_slug: str | None = None + request_id: str | None = None + + +class DataResourceFieldContract(ContractModel): + name: str + field_type: str = "string" + label: str | None = None + required: bool = False + editable: str = "always" + visible_on: list[str] = Field( + default_factory=lambda: ["list", "detail", "create", "edit"] + ) + description: str | None = None + related_resource: str | None = None + sensitive: bool = False + display_label: str | None = None + + +class DataResourceContract(ContractModel): + name: str + display_name: str + description: str = "" + fields: list[DataResourceFieldContract] = Field(default_factory=list) + read_only: bool = False + importable: bool = False + operations: dict[str, bool] = Field(default_factory=dict) + + +class AgentMethodFieldContract(ContractModel): + name: str + type: str | None = None + field_type: str = "CharField" + description: str | None = None + choices: list[Any] | None = None + default: Any = None + widget: str | None = None + required: bool = False + dynamic_choices: str | None = None + + +class AgentMethodContract(ContractModel): + name: str + method: str + params: dict[str, Any] | None = None + fields: list[AgentMethodFieldContract | dict[str, Any]] | None = None + description: str | None = None + nodes: dict[str, Any] | None = None + + +class AgentMethodsContract(ContractModel): + job_start: AgentMethodContract + job_stop: AgentMethodContract | None = None + job_status: AgentMethodContract | None = None + job_poll: AgentMethodContract | None = None + human_answer: AgentMethodContract | None = None + chat: AgentMethodContract | None = None + custom: dict[str, AgentMethodContract] | None = None + + +class ControllerContract(ContractModel): + """Canonical controller surface advertised by a Supervaizer server.""" + + controller_contract_version: str = Field(default=CONTROLLER_CONTRACT_VERSION) + api_base_path: str = Field(default=API_BASE_PATH) + endpoints: dict[str, str] = Field( + default_factory=lambda: { + endpoint.value: template + for endpoint, template in CONTROLLER_ENDPOINTS.items() + } + ) + + +class AgentRegistrationContract(ContractModel): + """Minimal schema for agent registration payloads consumed by Studio.""" + + id: str | None = None + slug: str + name: str + api_path: str + methods: AgentMethodsContract | dict[str, Any] = Field(default_factory=dict) + parameters_setup: list[dict[str, Any]] = Field(default_factory=list) + data_resources: list[DataResourceContract | dict[str, Any]] = Field( + default_factory=list + ) + + +class ServerRegistrationContract(ControllerContract): + """Minimal schema for server.register details.""" + + server_id: str + url: str + uri: str + api_version: str + environment: str | None = None + agents: list[AgentRegistrationContract] = Field(default_factory=list) + + +class JobStartRequest(ContractModel): + job_context: dict[str, Any] + job_fields: dict[str, Any] = Field(default_factory=dict) + encrypted_agent_parameters: str | None = None + + +class DynamicChoicesRequest(ContractModel): + workspace_id: str + mission_id: str + workspace_slug: str | None = None + + +class DynamicChoicesResponse(ContractModel): + choices: dict[str, list[Any]] = Field(default_factory=dict) + + +class CaseUpdateEvent(ContractModel): + name: str + payload: dict[str, Any] = Field(default_factory=dict) + cost: float = 0.0 + index: int | None = None + is_final: bool = False + + +class CaseUpdateRequest(ContractModel): + answer: dict[str, Any] + message: str | None = None + + +class DataResourceListResponse(ContractModel): + """Structured response shape for DataResource list operations.""" + + items: list[dict[str, Any]] = Field(default_factory=list) + + +def _endpoint_key(endpoint: ControllerEndpoint | str) -> str: + return endpoint.value if isinstance(endpoint, ControllerEndpoint) else endpoint + + +def resolve_controller_endpoint( + contract: ControllerContract | dict[str, Any], + endpoint: ControllerEndpoint | str, + **params: Any, +) -> str: + """Resolve a controller endpoint from a registered contract.""" + parsed = ( + contract + if isinstance(contract, ControllerContract) + else ControllerContract.model_validate(contract) + ) + endpoint_key = _endpoint_key(endpoint) + template = parsed.endpoints.get(endpoint_key) + if not template: + raise KeyError( + f"Controller endpoint {endpoint_key!r} is not advertised by the registered contract" + ) + try: + return template.format(**params) + except KeyError as exc: + missing = exc.args[0] + raise KeyError( + f"Missing parameter {missing!r} for controller endpoint {endpoint_key!r}" + ) from exc + + +def build_data_resource_context_headers( + *, + workspace_id: str | None = None, + workspace_slug: str | None = None, + mission_id: str | None = None, + request_id: str | None = None, +) -> dict[str, str]: + """Build Supervaize context headers for Studio DataResource proxy calls.""" + headers: dict[str, str] = {} + if workspace_id: + headers["X-Supervaize-Workspace-Id"] = str(workspace_id) + if workspace_slug: + headers["X-Supervaize-Workspace-Slug"] = str(workspace_slug) + if mission_id: + headers["X-Supervaize-Mission-Id"] = str(mission_id) + if request_id: + headers["X-Supervaize-Request-Id"] = str(request_id) + return headers + + +def controller_contract_info() -> dict[str, Any]: + """Return the JSON-serializable controller contract.""" + return ControllerContract().model_dump(mode="json") diff --git a/src/supervaizer/data_resource.py b/src/supervaizer/data_resource.py index 6655efc..3ac5c7e 100644 --- a/src/supervaizer/data_resource.py +++ b/src/supervaizer/data_resource.py @@ -51,6 +51,16 @@ class Editable(StrEnum): NEVER = "never" # Agent-controlled; never shown in a form input +class DataResourceContext(SvBaseModel): + """Studio request context passed to DataResource callbacks.""" + + workspace_id: str | None = None + workspace_slug: str | None = None + mission_id: str | None = None + agent_slug: str + request_id: str | None = None + + class DataResourceField(SvBaseModel): """Describes a single field in a DataResource for Studio rendering.""" @@ -75,6 +85,10 @@ class DataResourceField(SvBaseModel): default=None, description="Name of another DataResource this field FK-references", ) + sensitive: bool = Field( + default=False, + description="True when Studio should mask this field for non-manager users", + ) @property def display_label(self) -> str: @@ -120,22 +134,18 @@ class DataResource(SvBaseModel): read_only: bool = Field(default=False) importable: bool = Field(default=False, description="Enables CSV bulk import route") # Callbacks — excluded from model serialization - on_list: Callable[[], list[dict[str, Any]]] | None = Field( - default=None, exclude=True - ) - on_get: Callable[[str], dict[str, Any] | None] | None = Field( - default=None, exclude=True - ) - on_create: Callable[[dict[str, Any]], dict[str, Any]] | None = Field( + on_list: Callable[..., list[dict[str, Any]]] | None = Field( default=None, exclude=True ) - on_update: Callable[[str, dict[str, Any]], dict[str, Any] | None] | None = Field( + on_get: Callable[..., dict[str, Any] | None] | None = Field( default=None, exclude=True ) - on_delete: Callable[[str], bool] | None = Field(default=None, exclude=True) - on_import: Callable[[list[dict[str, Any]]], dict[str, Any]] | None = Field( + on_create: Callable[..., dict[str, Any]] | None = Field(default=None, exclude=True) + on_update: Callable[..., dict[str, Any] | None] | None = Field( default=None, exclude=True ) + on_delete: Callable[..., bool] | None = Field(default=None, exclude=True) + on_import: Callable[..., dict[str, Any]] | None = Field(default=None, exclude=True) @model_validator(mode="after") def check_callbacks(self) -> "DataResource": diff --git a/src/supervaizer/data_routes.py b/src/supervaizer/data_routes.py index ada711c..d5c26df 100644 --- a/src/supervaizer/data_routes.py +++ b/src/supervaizer/data_routes.py @@ -26,6 +26,7 @@ from __future__ import annotations +import inspect from typing import TYPE_CHECKING, Any from fastapi import ( @@ -34,12 +35,13 @@ Depends, HTTPException, Query, + Request, ) # <-- MODIFIED: removed Security, added Depends from fastapi.responses import JSONResponse from supervaizer.access import require_scope # <-- ADDED from supervaizer.common import log -from supervaizer.data_resource import DataResource +from supervaizer.data_resource import DataResource, DataResourceContext if TYPE_CHECKING: from supervaizer.agent import Agent @@ -76,7 +78,7 @@ def _add_resource_routes( op_id = _data_resource_operation_id(agent_slug, resource.name, "list") router.add_api_route( f"{prefix}/", - _make_list_handler(resource, prefix), + _make_list_handler(resource, prefix, agent_slug), methods=["GET"], # <-- REMOVED: Security(server.verify_api_key); api_router handles auth summary=f"List {resource.display_name_resolved}", @@ -88,7 +90,7 @@ def _add_resource_routes( op_id = _data_resource_operation_id(agent_slug, resource.name, "get") router.add_api_route( f"{prefix}/{{item_id}}", - _make_get_handler(resource, prefix), + _make_get_handler(resource, prefix, agent_slug), methods=["GET"], # <-- REMOVED: Security(server.verify_api_key); api_router handles auth summary=f"Get {resource.display_name_resolved}", @@ -100,7 +102,7 @@ def _add_resource_routes( op_id = _data_resource_operation_id(agent_slug, resource.name, "create") router.add_api_route( f"{prefix}/", - _make_create_handler(resource, prefix), + _make_create_handler(resource, prefix, agent_slug), methods=["POST"], dependencies=[ Depends(require_scope("write")) @@ -114,7 +116,7 @@ def _add_resource_routes( op_id = _data_resource_operation_id(agent_slug, resource.name, "update") router.add_api_route( f"{prefix}/{{item_id}}", - _make_update_handler(resource, prefix), + _make_update_handler(resource, prefix, agent_slug), methods=["PUT"], dependencies=[ Depends(require_scope("write")) @@ -128,7 +130,7 @@ def _add_resource_routes( op_id = _data_resource_operation_id(agent_slug, resource.name, "delete") router.add_api_route( f"{prefix}/{{item_id}}", - _make_delete_handler(resource, prefix), + _make_delete_handler(resource, prefix, agent_slug), methods=["DELETE"], dependencies=[ Depends(require_scope("write")) @@ -142,7 +144,7 @@ def _add_resource_routes( op_id = _data_resource_operation_id(agent_slug, resource.name, "import") router.add_api_route( f"{prefix}/import/", - _make_import_handler(resource, prefix), + _make_import_handler(resource, prefix, agent_slug), methods=["POST"], dependencies=[ Depends(require_scope("write")) @@ -153,22 +155,30 @@ def _add_resource_routes( ) -def _make_list_handler(r: DataResource, prefix: str) -> Any: +def _make_list_handler(r: DataResource, prefix: str, agent_slug: str) -> Any: async def _handler( + request: Request, skip: int = Query(default=0, ge=0), limit: int = Query(default=100, ge=1, le=1000), ) -> list[dict[str, Any]]: log.info(f"📥 GET {prefix}/ [DataResource list: {r.name}]") - result = r.on_list() # type: ignore[misc] + result = _call_with_context( + r.on_list, + _context_from_request(request, agent_slug), + ) return result[skip : skip + limit] return _handler -def _make_get_handler(r: DataResource, prefix: str) -> Any: - async def _handler(item_id: str) -> dict[str, Any]: +def _make_get_handler(r: DataResource, prefix: str, agent_slug: str) -> Any: + async def _handler(request: Request, item_id: str) -> dict[str, Any]: log.info(f"📥 GET {prefix}/{item_id} [DataResource get: {r.name}]") - result = r.on_get(item_id) # type: ignore[misc] + result = _call_with_context( + r.on_get, + _context_from_request(request, agent_slug), + item_id, + ) if result is None: raise HTTPException( status_code=404, detail=f"{r.name} '{item_id}' not found" @@ -178,10 +188,16 @@ async def _handler(item_id: str) -> dict[str, Any]: return _handler -def _make_create_handler(r: DataResource, prefix: str) -> Any: - async def _handler(data: dict[str, Any] = Body(...)) -> JSONResponse: +def _make_create_handler(r: DataResource, prefix: str, agent_slug: str) -> Any: + async def _handler( + request: Request, data: dict[str, Any] = Body(...) + ) -> JSONResponse: log.info(f"📥 POST {prefix}/ [DataResource create: {r.name}]") - result = r.on_create(data) # type: ignore[misc] + result = _call_with_context( + r.on_create, + _context_from_request(request, agent_slug), + data, + ) if not isinstance(result, dict) or "id" not in result: raise HTTPException( status_code=500, @@ -192,15 +208,20 @@ async def _handler(data: dict[str, Any] = Body(...)) -> JSONResponse: return _handler -def _make_update_handler(r: DataResource, prefix: str) -> Any: +def _make_update_handler(r: DataResource, prefix: str, agent_slug: str) -> Any: on_update = r.on_update assert on_update is not None # route registered only when on_update is set async def _handler( - item_id: str, data: dict[str, Any] = Body(...) + request: Request, item_id: str, data: dict[str, Any] = Body(...) ) -> dict[str, Any]: log.info(f"📥 PUT {prefix}/{item_id} [DataResource update: {r.name}]") - result = on_update(item_id, data) + result = _call_with_context( + on_update, + _context_from_request(request, agent_slug), + item_id, + data, + ) if result is None: raise HTTPException( status_code=404, detail=f"{r.name} '{item_id}' not found" @@ -210,10 +231,14 @@ async def _handler( return _handler -def _make_delete_handler(r: DataResource, prefix: str) -> Any: - async def _handler(item_id: str) -> JSONResponse: +def _make_delete_handler(r: DataResource, prefix: str, agent_slug: str) -> Any: + async def _handler(request: Request, item_id: str) -> JSONResponse: log.info(f"📥 DELETE {prefix}/{item_id} [DataResource delete: {r.name}]") - success = r.on_delete(item_id) # type: ignore[misc] + success = _call_with_context( + r.on_delete, + _context_from_request(request, agent_slug), + item_id, + ) if not success: raise HTTPException( status_code=404, detail=f"{r.name} '{item_id}' not found" @@ -223,9 +248,50 @@ async def _handler(item_id: str) -> JSONResponse: return _handler -def _make_import_handler(r: DataResource, prefix: str) -> Any: - async def _handler(records: list[dict[str, Any]] = Body(...)) -> dict[str, Any]: +def _make_import_handler(r: DataResource, prefix: str, agent_slug: str) -> Any: + async def _handler( + request: Request, records: list[dict[str, Any]] = Body(...) + ) -> dict[str, Any]: log.info(f"📥 POST {prefix}/import/ [DataResource import: {r.name}]") - return r.on_import(records) # type: ignore[misc] + return _call_with_context( + r.on_import, + _context_from_request(request, agent_slug), + records, + ) return _handler + + +def _context_from_request(request: Request, agent_slug: str) -> DataResourceContext: + return DataResourceContext( + workspace_id=request.headers.get("X-Supervaize-Workspace-Id"), + workspace_slug=request.headers.get("X-Supervaize-Workspace-Slug"), + mission_id=request.headers.get("X-Supervaize-Mission-Id"), + agent_slug=agent_slug, + request_id=request.headers.get("X-Supervaize-Request-Id"), + ) + + +def _accepts_context(callback: Any) -> bool: + try: + signature = inspect.signature(callback) + except (TypeError, ValueError): + return False + return "context" in signature.parameters + + +def _call_with_context( + callback: Any, + context: DataResourceContext, + *args: Any, +) -> Any: + if callback is None: + raise HTTPException( + status_code=501, detail="DataResource callback not configured" + ) + try: + if _accepts_context(callback): + return callback(*args, context=context) + return callback(*args) + except PermissionError as exc: + raise HTTPException(status_code=403, detail=str(exc)) from exc diff --git a/src/supervaizer/event.py b/src/supervaizer/event.py index df0b83c..f523a4d 100644 --- a/src/supervaizer/event.py +++ b/src/supervaizer/event.py @@ -10,11 +10,11 @@ # If a copy of the MPL was not distributed with this file, you can obtain one at # https://mozilla.org/MPL/2.0/. -from enum import Enum from typing import TYPE_CHECKING, Any, ClassVar, Dict from supervaizer.__version__ import VERSION from supervaizer.common import SvBaseModel +from supervaizer.contracts import EventType from supervaizer.lifecycle import EntityStatus if TYPE_CHECKING: @@ -24,24 +24,6 @@ from supervaizer.server import Server -class EventType(str, Enum): - AGENT_REGISTER = "agent.register" - SERVER_REGISTER = "server.register" - AGENT_WAKEUP = "agent.wakeup" - AGENT_SEND_ANOMALY = "agent.anomaly" - INTERMEDIARY = "agent.intermediary" - JOB_START_CONFIRMATION = "agent.job.start.confirmation" - JOB_END = "agent.job.end" - JOB_STATUS = "agent.job.status" - JOB_RESULT = "agent.job.result" - JOB_ERROR = "agent.job.error" - CASE_START = "agent.case.start" - CASE_END = "agent.case.end" - CASE_STATUS = "agent.case.status" - CASE_RESULT = "agent.case.result" - CASE_UPDATE = "agent.case.update" - - class AbstractEvent(SvBaseModel): supervaizer_VERSION: ClassVar[str] = VERSION source: Dict[str, Any] diff --git a/src/supervaizer/job_service.py b/src/supervaizer/job_service.py index 9767496..b6bb8c4 100644 --- a/src/supervaizer/job_service.py +++ b/src/supervaizer/job_service.py @@ -89,7 +89,13 @@ def service_job_finished(job: Job, server: "Server") -> None: Tested in tests/test_job_service.py """ - assert server.supervisor_account is not None, "No account defined" + if server.supervisor_account is None: + log.warning( + "[service_job_finished] No supervisor account defined for server, " + "skipping finished event for job {}", + job.id, + ) + return account = server.supervisor_account event = JobFinishedEvent( job=job, diff --git a/src/supervaizer/routes.py b/src/supervaizer/routes.py index e9aece9..cd5404f 100644 --- a/src/supervaizer/routes.py +++ b/src/supervaizer/routes.py @@ -43,6 +43,7 @@ ) from supervaizer.case import CaseNodeUpdate, Cases from supervaizer.common import SvBaseModel, log +from supervaizer.contracts import controller_contract_info from supervaizer.job import Job, JobContext, JobResponse, Jobs from supervaizer.job_service import service_job_custom, service_job_start from supervaizer.lifecycle import EntityStatus @@ -129,6 +130,11 @@ def create_default_routes(server: "Server") -> APIRouter: """Create default routes for the server.""" router = APIRouter(prefix="/supervaizer", tags=["Supervision"]) + @router.get("/contract", response_model=dict[str, Any]) + async def get_controller_contract() -> dict[str, Any]: + """Return the controller contract Studio should use for route resolution.""" + return controller_contract_info() + @router.get( "/jobs/{job_id}", response_model=JobResponse, diff --git a/src/supervaizer/server.py b/src/supervaizer/server.py index af652df..c0652a5 100644 --- a/src/supervaizer/server.py +++ b/src/supervaizer/server.py @@ -47,6 +47,7 @@ is_local_mode, log, ) +from supervaizer.contracts import controller_contract_info from supervaizer.instructions import display_instructions from supervaizer.routes import get_server # <-- MODIFIED: removed per-router imports from supervaizer.routers import ( @@ -630,11 +631,13 @@ def uri(self) -> str: def registration_info(self) -> Dict[str, Any]: """Get registration info for the server.""" assert self.public_key is not None, "Public key not initialized" + contract = controller_contract_info() return { "server_id": self.server_id, "url": self.public_url, "uri": self.uri, "api_version": API_VERSION, + **contract, "environment": self.environment, "public_key": str( self.public_key.public_bytes( diff --git a/tests/test_case.py b/tests/test_case.py index d5cddbe..7d9d6b5 100644 --- a/tests/test_case.py +++ b/tests/test_case.py @@ -290,6 +290,28 @@ def test_case_nodes_get_method_with_empty_list() -> None: assert node is None +def test_case_nodes_node_index(case_nodes_all_steps: CaseNodes) -> None: + assert case_nodes_all_steps.node_index("confirm_call") == 1 + assert case_nodes_all_steps.node_index("start_call", start=2) == 3 + + +def test_case_nodes_make_update(case_nodes_all_steps: CaseNodes) -> None: + update = case_nodes_all_steps.make_update( + "call_completed", + payload={"status": "done"}, + cost=1.5, + is_final=True, + upsert_to="start_call", + ) + + assert update.name == "call_completed" + assert update.payload == {"status": "done"} + assert update.cost == 1.5 + assert update.is_final is True + assert update.index == case_nodes_all_steps.node_index("start_call") + assert update.upsert is True + + def test_case_node_different_types() -> None: """Test CaseNode instantiation with all different CaseNodeType values.""" types_to_test = [ diff --git a/tests/test_contracts.py b/tests/test_contracts.py new file mode 100644 index 0000000..46b2d62 --- /dev/null +++ b/tests/test_contracts.py @@ -0,0 +1,103 @@ +# Copyright (c) 2024-2026 Alain Prasquier - Supervaize.com. All rights reserved. +# +# This Source Code Form is subject to the terms of the Mozilla Public License, v. 2.0. +# If a copy of the MPL was not distributed with this file, you can obtain one at +# https://mozilla.org/MPL/2.0/. + +"""Tests for versioned Studio controller contracts.""" + +from __future__ import annotations + +import importlib +import sys + +from supervaizer.contracts import ( + ControllerEndpoint, + ControllerContract, + ServerRegistrationContract, + build_data_resource_context_headers, + controller_contract_info, + resolve_controller_endpoint, +) + + +def test_controller_contract_endpoints_are_api_prefixed() -> None: + info = controller_contract_info() + + assert info["controller_contract_version"] == "1.0" + assert info["api_base_path"] == "/api" + assert ( + info["endpoints"]["POST_AGENT_JOB_START"] + == "/api/supervaizer/agents/{agent_slug}/jobs" + ) + assert ( + info["endpoints"]["DATA_RESOURCE"] + == "/api/agents/{agent_slug}/data/{resource_name}/" + ) + + +def test_contract_models_export_json_schema() -> None: + controller_schema = ControllerContract.model_json_schema() + server_schema = ServerRegistrationContract.model_json_schema() + + assert controller_schema["properties"]["endpoints"]["type"] == "object" + assert server_schema["properties"]["agents"]["type"] == "array" + + +def test_contract_module_import_does_not_load_controller_runtime() -> None: + sys.modules.pop("supervaizer.server", None) + sys.modules.pop("supervaizer.routes", None) + importlib.import_module("supervaizer.contracts") + + assert "supervaizer.server" not in sys.modules + assert "supervaizer.routes" not in sys.modules + + +def test_resolve_controller_endpoint() -> None: + contract = ControllerContract() + + assert ( + resolve_controller_endpoint( + contract, + ControllerEndpoint.POST_AGENT_JOB_START, + agent_slug="agent-interviewer", + ) + == "/api/supervaizer/agents/agent-interviewer/jobs" + ) + assert ( + resolve_controller_endpoint( + contract.model_dump(mode="json"), + "DATA_RESOURCE_ITEM", + agent_slug="agent-interviewer", + resource_name="contacts", + item_id="c1", + ) + == "/api/agents/agent-interviewer/data/contacts/c1" + ) + + +def test_resolve_controller_endpoint_rejects_unknown_endpoint() -> None: + contract = ControllerContract() + + try: + resolve_controller_endpoint(contract, "UNKNOWN") + except KeyError as exc: + assert "UNKNOWN" in str(exc) + else: + raise AssertionError("Expected unknown endpoint to raise KeyError") + + +def test_data_resource_context_headers() -> None: + headers = build_data_resource_context_headers( + workspace_id="1", + workspace_slug="team-slug", + mission_id="mission-1", + request_id="request-1", + ) + + assert headers == { + "X-Supervaize-Workspace-Id": "1", + "X-Supervaize-Workspace-Slug": "team-slug", + "X-Supervaize-Mission-Id": "mission-1", + "X-Supervaize-Request-Id": "request-1", + } diff --git a/tests/test_job_service.py b/tests/test_job_service.py index 5d5ce29..608c043 100644 --- a/tests/test_job_service.py +++ b/tests/test_job_service.py @@ -243,16 +243,22 @@ def test_service_job_finished(server_fixture: "Server", mocker: MockerFixture) - def test_service_job_finished_without_account( server_fixture: "Server", mocker: MockerFixture ) -> None: - """Test service_job_finished raises assertion error when no account is defined.""" + """Test service_job_finished is a no-op when no account is defined.""" # Create a mock job mock_job = mocker.MagicMock(spec=Job) + mock_job.id = str(uuid.uuid4()) + mock_event_class = mocker.patch("supervaizer.job_service.JobFinishedEvent") + mock_send_event = mocker.patch( + "supervaizer.account_service.send_event", return_value=None + ) # Remove supervisor account server_fixture.supervisor_account = None - # Should raise AssertionError - with pytest.raises(AssertionError, match="No account defined"): - service_job_finished(job=mock_job, server=server_fixture) + service_job_finished(job=mock_job, server=server_fixture) + + mock_event_class.assert_not_called() + mock_send_event.assert_not_called() @pytest.mark.asyncio diff --git a/tests/test_routes.py b/tests/test_routes.py index e49635d..17137d4 100644 --- a/tests/test_routes.py +++ b/tests/test_routes.py @@ -23,7 +23,7 @@ Job, Server, ) -from supervaizer.data_resource import DataResource +from supervaizer.data_resource import DataResource, DataResourceContext from supervaizer.parameter import ParametersSetup from supervaizer.routes import ( create_agents_routes, @@ -363,3 +363,63 @@ def test_data_resource_openapi_operation_ids_unique_per_agent( assert len(set(list_ids)) == 2 assert f"{agent_a.slug}_items_list" in list_ids assert f"{agent_b.slug}_items_list" in list_ids + + +def test_data_resource_callbacks_receive_context( + account_fixture: Account, + agent_method_fixture: AgentMethod, + parameters_setup_fixture: ParametersSetup, +) -> None: + captured: dict[str, DataResourceContext] = {} + + def on_list(*, context: DataResourceContext) -> list[dict[str, Any]]: + captured["context"] = context + return [{"id": "1"}] + + resource = DataResource(name="items", fields=[], on_list=on_list, read_only=True) + methods = AgentMethods(job_start=agent_method_fixture) + agent = Agent( + name="Context Agent", + author="a", + developer="d", + version="1.0.0", + description="d", + methods=methods, + parameters_setup=parameters_setup_fixture, + data_resources=[resource], + ) + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + server = Server( + scheme="http", + host="localhost", + port=8001, + environment="test", + mac_addr="E2-AC-ED-22-BF-B2", + debug=True, + agent_timeout=10, + private_key=private_key, + a2a_endpoints=False, + supervisor_account=account_fixture, + agents=[agent], + api_key="test-api-key", + ) + client = TestClient(server.app) + + response = client.get( + f"/api/agents/{agent.slug}/data/items/", + headers={ + "X-API-Key": "test-api-key", + "X-Supervaize-Workspace-Id": "team-1", + "X-Supervaize-Workspace-Slug": "team-slug", + "X-Supervaize-Mission-Id": "mission-1", + "X-Supervaize-Request-Id": "request-1", + }, + ) + + assert response.status_code == 200 + context = captured["context"] + assert context.agent_slug == agent.slug + assert context.workspace_id == "team-1" + assert context.workspace_slug == "team-slug" + assert context.mission_id == "mission-1" + assert context.request_id == "request-1" diff --git a/tests/test_server.py b/tests/test_server.py index 12b9c82..5cc5bd9 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -555,6 +555,11 @@ def test_server_registration_info(server_fixture: Server) -> None: assert "agents" in registration_info assert "url" in registration_info assert "server_id" in registration_info + assert registration_info["controller_contract_version"] == "1.0" + assert registration_info["api_base_path"] == "/api" + assert registration_info["endpoints"]["POST_AGENT_JOB_START"].startswith( + "/api/supervaizer/" + ) # Check values match the fixture assert registration_info.pop("public_key").startswith("-----BEGIN PUBLIC KEY") @@ -575,6 +580,9 @@ def test_server_registration_info(server_fixture: Server) -> None: assert registration_info == { "uri": "server:E2-AC-ED-22-BF-B2", "api_version": "v1", + "controller_contract_version": "1.0", + "api_base_path": "/api", + "endpoints": registration_info["endpoints"], "environment": "test", "api_key": "test-api-key", }