a whole ass Thing over here
This commit is contained in:
+1
-1
@@ -1,5 +1,5 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "lda-w"
|
name = "lda-wf"
|
||||||
version = "0.0.1"
|
version = "0.0.1"
|
||||||
description = "Add your description here"
|
description = "Add your description here"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
|
|||||||
@@ -0,0 +1,131 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
from wf_authoring import NodeReturn, node
|
||||||
|
from wf_core import END, RuntimeContext, RunStatus
|
||||||
|
from wf_mcp import AuthRecord, ConnectionConfig, FileStore, RawWorkflowPlan, WfMcpService
|
||||||
|
|
||||||
|
|
||||||
|
class EchoInput(BaseModel):
|
||||||
|
text: str
|
||||||
|
|
||||||
|
|
||||||
|
class EchoOutput(BaseModel):
|
||||||
|
echoed: str
|
||||||
|
|
||||||
|
|
||||||
|
class FinalizeInput(BaseModel):
|
||||||
|
echoed: str
|
||||||
|
|
||||||
|
|
||||||
|
class FinalizeOutput(BaseModel):
|
||||||
|
result: str
|
||||||
|
|
||||||
|
|
||||||
|
@node()
|
||||||
|
async def echo_tool(payload: EchoInput, ctx: RuntimeContext) -> EchoOutput:
|
||||||
|
return EchoOutput(echoed=payload.text)
|
||||||
|
|
||||||
|
|
||||||
|
@node(outcomes=("done",))
|
||||||
|
def finalize_tool(
|
||||||
|
payload: FinalizeInput, ctx: RuntimeContext
|
||||||
|
) -> NodeReturn[FinalizeOutput]:
|
||||||
|
return NodeReturn(
|
||||||
|
outcome="done",
|
||||||
|
output=FinalizeOutput(result=f"final:{payload.echoed}"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _local_temp_root() -> Path:
|
||||||
|
root = Path("test-artifacts") / "wf_mcp_store"
|
||||||
|
root.mkdir(parents=True, exist_ok=True)
|
||||||
|
return root
|
||||||
|
|
||||||
|
|
||||||
|
def test_file_store_round_trips_auth() -> None:
|
||||||
|
store = FileStore(_local_temp_root() / "auth_store")
|
||||||
|
record = AuthRecord(
|
||||||
|
connection_id="demo.personal",
|
||||||
|
scheme="oauth",
|
||||||
|
payload={"token": "secret"},
|
||||||
|
)
|
||||||
|
|
||||||
|
store.save_auth(record)
|
||||||
|
loaded = store.load_auth("demo.personal")
|
||||||
|
|
||||||
|
assert loaded == record
|
||||||
|
|
||||||
|
|
||||||
|
def test_service_builds_namespaced_catalog() -> None:
|
||||||
|
service = WfMcpService(store=FileStore(_local_temp_root() / "catalog_store"))
|
||||||
|
service.register_connection(
|
||||||
|
ConnectionConfig(id="demo.personal", server="demo", account="personal")
|
||||||
|
)
|
||||||
|
service.register_specs("demo.personal", echo_tool, finalize_tool)
|
||||||
|
|
||||||
|
payload = service.get_catalog().as_payload()
|
||||||
|
names = [node["qualified_name"] for node in payload["nodes"]]
|
||||||
|
|
||||||
|
assert names == [
|
||||||
|
"demo.personal.echo_tool",
|
||||||
|
"demo.personal.finalize_tool",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_service_compiles_and_runs_raw_plan() -> None:
|
||||||
|
service = WfMcpService(store=FileStore(_local_temp_root() / "run_store"))
|
||||||
|
service.register_connection(
|
||||||
|
ConnectionConfig(id="demo.personal", server="demo", account="personal")
|
||||||
|
)
|
||||||
|
service.register_specs("demo.personal", echo_tool, finalize_tool)
|
||||||
|
|
||||||
|
plan = RawWorkflowPlan(
|
||||||
|
name="demo_plan",
|
||||||
|
input_schema={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"text": {"type": "string"}},
|
||||||
|
"required": ["text"],
|
||||||
|
},
|
||||||
|
state_schema={
|
||||||
|
"fields": {
|
||||||
|
"echoed": {"type": "string"},
|
||||||
|
"result": {"type": "string"},
|
||||||
|
}
|
||||||
|
},
|
||||||
|
output_schema={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"result": {"type": "string"}},
|
||||||
|
"required": ["result"],
|
||||||
|
},
|
||||||
|
start="echo",
|
||||||
|
nodes=[
|
||||||
|
{
|
||||||
|
"id": "echo",
|
||||||
|
"type": "node",
|
||||||
|
"node": "demo.personal.echo_tool",
|
||||||
|
"in_map": {"input.text": "text"},
|
||||||
|
"out_map": {"echoed": "state.echoed"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "finalize",
|
||||||
|
"type": "node",
|
||||||
|
"node": "demo.personal.finalize_tool",
|
||||||
|
"in_map": {"state.echoed": "echoed"},
|
||||||
|
"out_map": {"result": "state.result"},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
edges=[
|
||||||
|
{"from": "echo", "outcome": "ok", "to": "finalize"},
|
||||||
|
{"from": "finalize", "outcome": "done", "to": END},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
run = asyncio.run(service.run_workflow_from_plan(plan, {"text": "hello"}))
|
||||||
|
|
||||||
|
assert run.status == RunStatus.COMPLETED
|
||||||
|
assert run.output == {"result": "final:hello"}
|
||||||
@@ -249,7 +249,7 @@ wheels = [
|
|||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "lda-w"
|
name = "lda-wf"
|
||||||
version = "0.0.1"
|
version = "0.0.1"
|
||||||
source = { virtual = "." }
|
source = { virtual = "." }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
|||||||
+38
-4
@@ -3,7 +3,7 @@ from __future__ import annotations
|
|||||||
from inspect import Parameter, iscoroutinefunction, signature
|
from inspect import Parameter, iscoroutinefunction, signature
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Generic, TypeVar, cast, get_args, get_origin, get_type_hints, overload
|
from typing import Any, Generic, Literal, TypeVar, cast, get_args, get_origin, get_type_hints, overload
|
||||||
|
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
@@ -31,6 +31,10 @@ class NodeReturn(Generic[OutputT]):
|
|||||||
output: OutputT
|
output: OutputT
|
||||||
|
|
||||||
|
|
||||||
|
def _default_outcome(spec: "NodeSpec[Any, Any]") -> str:
|
||||||
|
return spec.outcomes[0]
|
||||||
|
|
||||||
|
|
||||||
def _coerce_registry_result(
|
def _coerce_registry_result(
|
||||||
*,
|
*,
|
||||||
node_name: str,
|
node_name: str,
|
||||||
@@ -151,7 +155,7 @@ class NodeSpec(Generic[InputT, OutputT]):
|
|||||||
return _coerce_registry_result(
|
return _coerce_registry_result(
|
||||||
node_name=self.name,
|
node_name=self.name,
|
||||||
output_model=self.output_model,
|
output_model=self.output_model,
|
||||||
default_outcome=self.outcomes[0],
|
default_outcome=_default_outcome(self),
|
||||||
raw=cast(NodeReturn[BaseModel] | BaseModel, raw),
|
raw=cast(NodeReturn[BaseModel] | BaseModel, raw),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -174,7 +178,7 @@ class NodeSpec(Generic[InputT, OutputT]):
|
|||||||
return _coerce_registry_result(
|
return _coerce_registry_result(
|
||||||
node_name=self.name,
|
node_name=self.name,
|
||||||
output_model=self.output_model,
|
output_model=self.output_model,
|
||||||
default_outcome=self.outcomes[0],
|
default_outcome=_default_outcome(self),
|
||||||
raw=cast(NodeReturn[BaseModel] | BaseModel, raw),
|
raw=cast(NodeReturn[BaseModel] | BaseModel, raw),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -258,10 +262,40 @@ def node(
|
|||||||
def build_registry(
|
def build_registry(
|
||||||
*specs: NodeSpec[Any, Any],
|
*specs: NodeSpec[Any, Any],
|
||||||
) -> dict[str, SyncRegistryHandler]:
|
) -> dict[str, SyncRegistryHandler]:
|
||||||
return {spec.name: spec.to_registry_handler() for spec in specs}
|
return _build_registry(specs, export="sync")
|
||||||
|
|
||||||
|
|
||||||
def build_async_registry(
|
def build_async_registry(
|
||||||
*specs: NodeSpec[Any, Any],
|
*specs: NodeSpec[Any, Any],
|
||||||
) -> dict[str, AsyncRegistryHandler]:
|
) -> dict[str, AsyncRegistryHandler]:
|
||||||
|
return _build_registry(specs, export="async")
|
||||||
|
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def _build_registry(
|
||||||
|
specs: tuple[NodeSpec[Any, Any], ...],
|
||||||
|
*,
|
||||||
|
export: Literal["sync"],
|
||||||
|
) -> dict[str, SyncRegistryHandler]:
|
||||||
|
...
|
||||||
|
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def _build_registry(
|
||||||
|
specs: tuple[NodeSpec[Any, Any], ...],
|
||||||
|
*,
|
||||||
|
export: Literal["async"],
|
||||||
|
) -> dict[str, AsyncRegistryHandler]:
|
||||||
|
...
|
||||||
|
|
||||||
|
|
||||||
|
def _build_registry(
|
||||||
|
specs: tuple[NodeSpec[Any, Any], ...],
|
||||||
|
*,
|
||||||
|
export: Literal["sync", "async"],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
if export == "sync":
|
||||||
|
return {spec.name: spec.to_registry_handler() for spec in specs}
|
||||||
|
if export == "async":
|
||||||
return {spec.name: spec.to_async_registry_handler() for spec in specs}
|
return {spec.name: spec.to_async_registry_handler() for spec in specs}
|
||||||
|
raise ValueError(f"unknown registry export mode {export!r}")
|
||||||
|
|||||||
@@ -0,0 +1,26 @@
|
|||||||
|
from .catalog import CombinedCatalog
|
||||||
|
from .connections import ConnectionRegistry, parse_connection_id, qualify_node_name
|
||||||
|
from .models import (
|
||||||
|
AuthRecord,
|
||||||
|
CatalogNodeEntry,
|
||||||
|
CatalogSnapshot,
|
||||||
|
ConnectionConfig,
|
||||||
|
RawWorkflowPlan,
|
||||||
|
)
|
||||||
|
from .service import WfMcpService
|
||||||
|
from .store import FileStore, Store
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"AuthRecord",
|
||||||
|
"CatalogNodeEntry",
|
||||||
|
"CatalogSnapshot",
|
||||||
|
"CombinedCatalog",
|
||||||
|
"ConnectionConfig",
|
||||||
|
"ConnectionRegistry",
|
||||||
|
"FileStore",
|
||||||
|
"RawWorkflowPlan",
|
||||||
|
"Store",
|
||||||
|
"WfMcpService",
|
||||||
|
"parse_connection_id",
|
||||||
|
"qualify_node_name",
|
||||||
|
]
|
||||||
@@ -0,0 +1,66 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from wf_authoring import NodeCatalog, NodeSpec
|
||||||
|
|
||||||
|
from .connections import qualify_node_name
|
||||||
|
from .models import CatalogNodeEntry, CatalogSnapshot
|
||||||
|
|
||||||
|
|
||||||
|
def snapshot_from_specs(
|
||||||
|
connection_id: str,
|
||||||
|
*,
|
||||||
|
specs: dict[str, NodeSpec[Any, Any]],
|
||||||
|
fetched_at_epoch_ms: int,
|
||||||
|
max_age_seconds: int,
|
||||||
|
) -> CatalogSnapshot:
|
||||||
|
catalog = NodeCatalog.from_specs(*specs.values())
|
||||||
|
nodes = [
|
||||||
|
CatalogNodeEntry(
|
||||||
|
qualified_name=entry.name
|
||||||
|
if entry.name.startswith(f"{connection_id}.")
|
||||||
|
else qualify_node_name(connection_id, entry.name),
|
||||||
|
connection_id=connection_id,
|
||||||
|
local_name=entry.name.removeprefix(f"{connection_id}."),
|
||||||
|
description=entry.description,
|
||||||
|
outcomes=entry.outcomes,
|
||||||
|
input_schema=entry.input_schema,
|
||||||
|
output_schema=entry.output_schema,
|
||||||
|
)
|
||||||
|
for entry in catalog.entries()
|
||||||
|
]
|
||||||
|
return CatalogSnapshot(
|
||||||
|
connection_id=connection_id,
|
||||||
|
fetched_at_epoch_ms=fetched_at_epoch_ms,
|
||||||
|
max_age_seconds=max_age_seconds,
|
||||||
|
nodes=nodes,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class CombinedCatalog:
|
||||||
|
snapshots: dict[str, CatalogSnapshot] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def entries(self) -> list[CatalogNodeEntry]:
|
||||||
|
result: list[CatalogNodeEntry] = []
|
||||||
|
for snapshot in self.snapshots.values():
|
||||||
|
result.extend(snapshot.nodes)
|
||||||
|
return sorted(result, key=lambda entry: entry.qualified_name)
|
||||||
|
|
||||||
|
def as_payload(self) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"nodes": [
|
||||||
|
{
|
||||||
|
"qualified_name": entry.qualified_name,
|
||||||
|
"connection_id": entry.connection_id,
|
||||||
|
"local_name": entry.local_name,
|
||||||
|
"description": entry.description,
|
||||||
|
"outcomes": list(entry.outcomes),
|
||||||
|
"input_schema": entry.input_schema,
|
||||||
|
"output_schema": entry.output_schema,
|
||||||
|
}
|
||||||
|
for entry in self.entries()
|
||||||
|
]
|
||||||
|
}
|
||||||
@@ -0,0 +1,41 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
|
||||||
|
from .models import ConnectionConfig
|
||||||
|
|
||||||
|
|
||||||
|
def parse_connection_id(connection_id: str) -> tuple[str, str]:
|
||||||
|
if "." not in connection_id:
|
||||||
|
raise ValueError(
|
||||||
|
"connection id must look like '<server>.<account>'"
|
||||||
|
)
|
||||||
|
server, account = connection_id.split(".", 1)
|
||||||
|
if not server or not account:
|
||||||
|
raise ValueError(
|
||||||
|
"connection id must look like '<server>.<account>'"
|
||||||
|
)
|
||||||
|
return server, account
|
||||||
|
|
||||||
|
|
||||||
|
def qualify_node_name(connection_id: str, local_name: str) -> str:
|
||||||
|
parse_connection_id(connection_id)
|
||||||
|
if not local_name:
|
||||||
|
raise ValueError("local node name must not be empty")
|
||||||
|
return f"{connection_id}.{local_name}"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class ConnectionRegistry:
|
||||||
|
connections: dict[str, ConnectionConfig] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def register(self, connection: ConnectionConfig) -> None:
|
||||||
|
parse_connection_id(connection.id)
|
||||||
|
self.connections[connection.id] = connection
|
||||||
|
|
||||||
|
def get(self, connection_id: str) -> ConnectionConfig:
|
||||||
|
return self.connections[connection_id]
|
||||||
|
|
||||||
|
def list_enabled(self) -> list[ConnectionConfig]:
|
||||||
|
return [connection for connection in self.connections.values() if connection.enabled]
|
||||||
|
|
||||||
@@ -0,0 +1,64 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import asdict, dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class ConnectionConfig:
|
||||||
|
id: str
|
||||||
|
server: str
|
||||||
|
account: str
|
||||||
|
enabled: bool = True
|
||||||
|
metadata: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class AuthRecord:
|
||||||
|
connection_id: str
|
||||||
|
scheme: str
|
||||||
|
payload: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class CatalogNodeEntry:
|
||||||
|
qualified_name: str
|
||||||
|
connection_id: str
|
||||||
|
local_name: str
|
||||||
|
description: str | None
|
||||||
|
outcomes: tuple[str, ...]
|
||||||
|
input_schema: dict[str, Any]
|
||||||
|
output_schema: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class CatalogSnapshot:
|
||||||
|
connection_id: str
|
||||||
|
fetched_at_epoch_ms: int
|
||||||
|
max_age_seconds: int
|
||||||
|
nodes: list[CatalogNodeEntry] = field(default_factory=list)
|
||||||
|
|
||||||
|
def is_stale(self, now_epoch_ms: int) -> bool:
|
||||||
|
age_ms = now_epoch_ms - self.fetched_at_epoch_ms
|
||||||
|
return age_ms > self.max_age_seconds * 1000
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class RawWorkflowPlan:
|
||||||
|
name: str
|
||||||
|
input_schema: dict[str, Any]
|
||||||
|
state_schema: dict[str, Any]
|
||||||
|
output_schema: dict[str, Any]
|
||||||
|
start: str
|
||||||
|
nodes: list[dict[str, Any]]
|
||||||
|
edges: list[dict[str, Any]]
|
||||||
|
|
||||||
|
|
||||||
|
def dump_catalog_snapshot(snapshot: CatalogSnapshot) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"connection_id": snapshot.connection_id,
|
||||||
|
"fetched_at_epoch_ms": snapshot.fetched_at_epoch_ms,
|
||||||
|
"max_age_seconds": snapshot.max_age_seconds,
|
||||||
|
"nodes": [asdict(node) for node in snapshot.nodes],
|
||||||
|
}
|
||||||
|
|
||||||
@@ -0,0 +1,115 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from wf_authoring import NodeSpec, build_async_registry
|
||||||
|
from wf_core import NodeUse, Workflow, execute_workflow_async
|
||||||
|
|
||||||
|
from .catalog import CombinedCatalog, snapshot_from_specs
|
||||||
|
from .connections import ConnectionRegistry, parse_connection_id, qualify_node_name
|
||||||
|
from .models import AuthRecord, CatalogSnapshot, ConnectionConfig, RawWorkflowPlan
|
||||||
|
from .store import Store
|
||||||
|
|
||||||
|
|
||||||
|
def _qualify_spec(connection_id: str, spec: NodeSpec[Any, Any]) -> NodeSpec[Any, Any]:
|
||||||
|
return NodeSpec(
|
||||||
|
name=qualify_node_name(connection_id, spec.name),
|
||||||
|
input_model=spec.input_model,
|
||||||
|
output_model=spec.output_model,
|
||||||
|
outcomes=spec.outcomes,
|
||||||
|
fn=spec.fn,
|
||||||
|
description=spec.description,
|
||||||
|
is_async=spec.is_async,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class WfMcpService:
|
||||||
|
store: Store
|
||||||
|
default_catalog_max_age_seconds: int = 300
|
||||||
|
connections: ConnectionRegistry = field(default_factory=ConnectionRegistry)
|
||||||
|
specs_by_connection: dict[str, dict[str, NodeSpec[Any, Any]]] = field(
|
||||||
|
default_factory=dict
|
||||||
|
)
|
||||||
|
|
||||||
|
def register_connection(self, connection: ConnectionConfig) -> None:
|
||||||
|
parse_connection_id(connection.id)
|
||||||
|
self.connections.register(connection)
|
||||||
|
|
||||||
|
def save_auth(self, record: AuthRecord) -> None:
|
||||||
|
self.store.save_auth(record)
|
||||||
|
|
||||||
|
def load_auth(self, connection_id: str) -> AuthRecord | None:
|
||||||
|
return self.store.load_auth(connection_id)
|
||||||
|
|
||||||
|
def register_specs(
|
||||||
|
self,
|
||||||
|
connection_id: str,
|
||||||
|
*specs: NodeSpec[Any, Any],
|
||||||
|
max_age_seconds: int | None = None,
|
||||||
|
) -> None:
|
||||||
|
self.connections.get(connection_id)
|
||||||
|
qualified_specs = {
|
||||||
|
qualify_node_name(connection_id, spec.name): _qualify_spec(connection_id, spec)
|
||||||
|
for spec in specs
|
||||||
|
}
|
||||||
|
self.specs_by_connection[connection_id] = qualified_specs
|
||||||
|
snapshot = snapshot_from_specs(
|
||||||
|
connection_id,
|
||||||
|
specs=qualified_specs,
|
||||||
|
fetched_at_epoch_ms=int(time.time() * 1000),
|
||||||
|
max_age_seconds=max_age_seconds or self.default_catalog_max_age_seconds,
|
||||||
|
)
|
||||||
|
self.store.save_catalog(snapshot)
|
||||||
|
|
||||||
|
def get_catalog(self) -> CombinedCatalog:
|
||||||
|
snapshots: dict[str, CatalogSnapshot] = {}
|
||||||
|
for connection in self.connections.list_enabled():
|
||||||
|
snapshot = self.store.load_catalog(connection.id)
|
||||||
|
if snapshot is not None:
|
||||||
|
snapshots[connection.id] = snapshot
|
||||||
|
return CombinedCatalog(snapshots=snapshots)
|
||||||
|
|
||||||
|
def compile_plan(self, plan: RawWorkflowPlan) -> Workflow:
|
||||||
|
node_defs: dict[str, Any] = {}
|
||||||
|
for step in plan.nodes:
|
||||||
|
if step.get("type") != "node":
|
||||||
|
continue
|
||||||
|
qualified_name = step["node"]
|
||||||
|
spec = self._get_qualified_spec(qualified_name)
|
||||||
|
node_defs[qualified_name] = spec.to_node_def()
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"name": plan.name,
|
||||||
|
"input_schema": plan.input_schema,
|
||||||
|
"state_schema": plan.state_schema,
|
||||||
|
"output_schema": plan.output_schema,
|
||||||
|
"start": plan.start,
|
||||||
|
"node_defs": [node.model_dump() for node in node_defs.values()],
|
||||||
|
"nodes": plan.nodes,
|
||||||
|
"edges": plan.edges,
|
||||||
|
}
|
||||||
|
return Workflow.model_validate(payload)
|
||||||
|
|
||||||
|
async def run_workflow_from_plan(
|
||||||
|
self,
|
||||||
|
plan: RawWorkflowPlan,
|
||||||
|
workflow_input: dict[str, Any],
|
||||||
|
):
|
||||||
|
workflow = self.compile_plan(plan)
|
||||||
|
specs = [
|
||||||
|
self._get_qualified_spec(node.node)
|
||||||
|
for node in workflow.nodes
|
||||||
|
if isinstance(node, NodeUse)
|
||||||
|
]
|
||||||
|
registry = build_async_registry(*specs)
|
||||||
|
return await execute_workflow_async(workflow, workflow_input, registry)
|
||||||
|
|
||||||
|
def _get_qualified_spec(self, qualified_name: str) -> NodeSpec[Any, Any]:
|
||||||
|
connection_id, _ = qualified_name.rsplit(".", 1)
|
||||||
|
specs = self.specs_by_connection.get(connection_id)
|
||||||
|
if specs is None or qualified_name not in specs:
|
||||||
|
raise KeyError(f"unknown qualified node {qualified_name!r}")
|
||||||
|
return specs[qualified_name]
|
||||||
@@ -0,0 +1,81 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from dataclasses import asdict
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from .models import AuthRecord, CatalogNodeEntry, CatalogSnapshot
|
||||||
|
|
||||||
|
|
||||||
|
class Store:
|
||||||
|
def save_auth(self, record: AuthRecord) -> None:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def load_auth(self, connection_id: str) -> AuthRecord | None:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def save_catalog(self, snapshot: CatalogSnapshot) -> None:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def load_catalog(self, connection_id: str) -> CatalogSnapshot | None:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
|
class FileStore(Store):
|
||||||
|
def __init__(self, root: Path) -> None:
|
||||||
|
self.root = root
|
||||||
|
self.root.mkdir(parents=True, exist_ok=True)
|
||||||
|
self.auth_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
self.catalog_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def auth_dir(self) -> Path:
|
||||||
|
return self.root / "auth"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def catalog_dir(self) -> Path:
|
||||||
|
return self.root / "catalog"
|
||||||
|
|
||||||
|
def _auth_path(self, connection_id: str) -> Path:
|
||||||
|
return self.auth_dir / f"{connection_id}.json"
|
||||||
|
|
||||||
|
def _catalog_path(self, connection_id: str) -> Path:
|
||||||
|
return self.catalog_dir / f"{connection_id}.json"
|
||||||
|
|
||||||
|
def save_auth(self, record: AuthRecord) -> None:
|
||||||
|
self._auth_path(record.connection_id).write_text(
|
||||||
|
json.dumps(asdict(record), indent=2),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
|
||||||
|
def load_auth(self, connection_id: str) -> AuthRecord | None:
|
||||||
|
path = self._auth_path(connection_id)
|
||||||
|
if not path.exists():
|
||||||
|
return None
|
||||||
|
data = json.loads(path.read_text(encoding="utf-8"))
|
||||||
|
return AuthRecord(**data)
|
||||||
|
|
||||||
|
def save_catalog(self, snapshot: CatalogSnapshot) -> None:
|
||||||
|
payload = {
|
||||||
|
"connection_id": snapshot.connection_id,
|
||||||
|
"fetched_at_epoch_ms": snapshot.fetched_at_epoch_ms,
|
||||||
|
"max_age_seconds": snapshot.max_age_seconds,
|
||||||
|
"nodes": [asdict(node) for node in snapshot.nodes],
|
||||||
|
}
|
||||||
|
self._catalog_path(snapshot.connection_id).write_text(
|
||||||
|
json.dumps(payload, indent=2),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
|
||||||
|
def load_catalog(self, connection_id: str) -> CatalogSnapshot | None:
|
||||||
|
path = self._catalog_path(connection_id)
|
||||||
|
if not path.exists():
|
||||||
|
return None
|
||||||
|
data = json.loads(path.read_text(encoding="utf-8"))
|
||||||
|
return CatalogSnapshot(
|
||||||
|
connection_id=data["connection_id"],
|
||||||
|
fetched_at_epoch_ms=data["fetched_at_epoch_ms"],
|
||||||
|
max_age_seconds=data["max_age_seconds"],
|
||||||
|
nodes=[CatalogNodeEntry(**node) for node in data.get("nodes", [])],
|
||||||
|
)
|
||||||
|
|
||||||
Reference in New Issue
Block a user