Files
lda-wf/tests/wf_api/test_source_admin_api.py
T

298 lines
8.9 KiB
Python

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"
)
)