from __future__ import annotations import asyncio from typing import Any import pytest from wf_api import WorkflowSourceAdminApi, WorkflowSourceAdminSurface from wf_api.models import RawWorkflowPlan from wf_api.operation_context import WorkflowOperationContext from wf_api.saved_subgraphs import SavedSubgraphTree from wf_api.source_admin import WorkflowSourceDiagnosticsProvider from wf_artifacts import WorkflowArtifact, WorkflowDeployment from wf_authoring import NodeSpec from wf_core import RunState from wf_platform import ( CapabilityBuckets, CapabilitySource, SourcePermissions, SourceVisibility, ) class DummyEvents: def record_event(self, event: object) -> None: pass def record_workflow_event( self, event_type: str, *, capability_id: str, payload: dict[str, Any], ) -> None: pass class DummyRuntime: async def run_workflow_from_plan( self, plan: RawWorkflowPlan, workflow_input: dict[str, Any], deployment: WorkflowDeployment | None = None, artifact: WorkflowArtifact | None = None, saved_subgraph_tree: SavedSubgraphTree | None = None, ) -> RunState: raise AssertionError("source admin tests must not run workflows") async def resume_workflow_from_plan( self, plan: RawWorkflowPlan, run: RunState, *, resume_payload: dict[str, Any], resume_outcome: str, deployment: WorkflowDeployment | None = None, artifact: WorkflowArtifact | None = None, saved_subgraph_tree: SavedSubgraphTree | None = None, ) -> RunState: raise AssertionError("source admin tests must not resume workflows") class StaticSpecProvider: def __init__(self, sources: dict[str, CapabilitySource]) -> None: self._sources = sources @property def capability_sources(self) -> dict[str, CapabilitySource]: return self._sources def get_qualified_spec(self, qualified_name: str) -> NodeSpec[Any, Any]: raise KeyError(f"unknown capability {qualified_name!r}") def _api(*sources: CapabilitySource) -> WorkflowSourceAdminApi: provider = StaticSpecProvider({source.id: source for source in sources}) return WorkflowSourceAdminApi( WorkflowOperationContext( artifact_store=None, draft_workspace_store=None, run_store=None, events=DummyEvents(), specs=provider, runtime=DummyRuntime(), live_sources=None, ) ) def _source(source_id: str, *, enabled: bool = True) -> CapabilitySource: return CapabilitySource( id=source_id, kind="connection", enabled=enabled, capabilities=CapabilityBuckets(), visibility=SourceVisibility( planner=True, client=True, admin_dashboard=True, ), permissions=SourcePermissions(calls_upstream=True), description=f"{source_id} source", ) def test_source_admin_lists_compact_sources_in_id_order() -> None: api = _api(_source("zeta.personal"), _source("alpha.personal", enabled=False)) payload = asyncio.run(api.list_sources(limit=10)) assert payload["total"] == 2 assert payload["next_cursor"] is None assert [source["id"] for source in payload["sources"]] == [ "alpha.personal", "zeta.personal", ] assert payload["sources"][0]["enabled"] is False assert payload["sources"][1]["description"] == "zeta.personal source" def test_source_admin_pages_sources() -> None: api = _api(_source("a"), _source("b"), _source("c")) first = asyncio.run(api.list_sources(limit=2)) second = asyncio.run(api.list_sources(cursor=first["next_cursor"], limit=2)) assert [source["id"] for source in first["sources"]] == ["a", "b"] assert first["next_cursor"] == "2" assert [source["id"] for source in second["sources"]] == ["c"] assert second["next_cursor"] is None def test_source_admin_inspects_full_source_inventory() -> None: api = _api(_source("demo.personal")) payload = asyncio.run(api.inspect_source(source_id="demo.personal")) assert payload["id"] == "demo.personal" assert payload["kind"] == "connection" assert payload["description"] == "demo.personal source" assert payload["visibility"]["planner"] is True assert payload["permissions"]["calls_upstream"] is True def test_source_admin_inspect_unknown_source_raises_clear_key_error() -> None: api = _api(_source("demo.personal")) with pytest.raises(KeyError, match="unknown source 'missing.source'"): asyncio.run(api.inspect_source(source_id="missing.source")) def test_source_admin_api_satisfies_surface_protocol() -> None: api: WorkflowSourceAdminSurface = _api(_source("demo.personal")) assert api is not None class _Diagnostics: def __init__(self) -> None: self.calls: list[str] = [] def diagnose_source(self, source_id: str) -> dict[str, object]: self.calls.append(source_id) return { "source_id": source_id, "status": "ok", "auth": {"record_present": True}, "diagnostics": [], } def _api_with_diagnostics( *sources: CapabilitySource, diagnostics: WorkflowSourceDiagnosticsProvider | None = None, ) -> WorkflowSourceAdminApi: provider = StaticSpecProvider({source.id: source for source in sources}) return WorkflowSourceAdminApi( WorkflowOperationContext( artifact_store=None, draft_workspace_store=None, run_store=None, events=DummyEvents(), specs=provider, runtime=DummyRuntime(), live_sources=None, ), diagnostics=diagnostics, ) def test_inspect_source_includes_optional_diagnostics() -> None: provider = _Diagnostics() payload = asyncio.run( _api_with_diagnostics( _source("demo.personal"), diagnostics=provider, ).inspect_source(source_id="demo.personal") ) assert payload["id"] == "demo.personal" diagnostics = payload.get("diagnostics") assert diagnostics is not None assert diagnostics.get("source_id") == "demo.personal" assert provider.calls == ["demo.personal"] def test_inspect_source_omits_diagnostics_without_provider() -> None: payload = asyncio.run( _api_with_diagnostics(_source("demo.personal")).inspect_source( source_id="demo.personal" ) ) assert payload["id"] == "demo.personal" assert "diagnostics" not in payload class _BrokenDiagnostics: def diagnose_source(self, source_id: str) -> dict[str, object]: raise RuntimeError("diagnostics exploded") def test_inspect_source_tolerates_diagnostics_provider_failure() -> None: payload = asyncio.run( _api_with_diagnostics( _source("demo.personal"), diagnostics=_BrokenDiagnostics(), ).inspect_source(source_id="demo.personal") ) assert payload["id"] == "demo.personal" diagnostics = payload.get("diagnostics") assert diagnostics is not None assert diagnostics["status"] == "error" message = diagnostics.get("message") assert message is not None assert "Diagnostics unavailable" in message def test_diagnose_source_uses_provider() -> None: payload = asyncio.run( _api_with_diagnostics( _source("demo.personal"), diagnostics=_Diagnostics(), ).diagnose_source(source_id="demo.personal") ) assert payload["status"] == "ok" auth = payload.get("auth") assert auth is not None assert auth.get("record_present") is True class _ExtendedDiagnostics: def diagnose_source(self, source_id: str) -> dict[str, object]: return { "source_id": source_id, "status": "degraded", "diagnostics": [], "provider_latency_ms": 12, } def test_diagnose_source_preserves_provider_extensions() -> None: payload = asyncio.run( _api_with_diagnostics( _source("demo.personal"), diagnostics=_ExtendedDiagnostics(), ).diagnose_source(source_id="demo.personal") ) assert payload["status"] == "degraded" assert dict(payload)["provider_latency_ms"] == 12 def test_diagnose_source_without_provider_returns_basic_status() -> None: payload = asyncio.run( _api_with_diagnostics(_source("demo.personal")).diagnose_source( source_id="demo.personal" ) ) assert payload == { "source_id": "demo.personal", "status": "unknown", "diagnostics": [], "message": "No source diagnostics provider is configured.", } def test_diagnose_source_unknown_raises_key_error() -> None: with pytest.raises(KeyError, match="unknown source 'missing.source'"): asyncio.run( _api_with_diagnostics(_source("demo.personal")).diagnose_source( source_id="missing.source" ) )