convert Callables to real protocols

ts the stuff langgraph do
This commit is contained in:
lda
2026-05-06 22:07:23 +07:00 Verified
parent 93fe9da816
commit 05a0db63c8
+40 -22
View File
@@ -7,6 +7,7 @@ from typing import (
Any, Any,
Generic, Generic,
Literal, Literal,
Protocol,
TypeVar, TypeVar,
cast, cast,
get_args, get_args,
@@ -22,28 +23,45 @@ from wf_core import NodeDef, RuntimeContext, SchemaRef
InputT = TypeVar("InputT", bound=BaseModel) InputT = TypeVar("InputT", bound=BaseModel)
OutputT = TypeVar("OutputT", bound=BaseModel) OutputT = TypeVar("OutputT", bound=BaseModel)
ContextNodeCallable = Callable[
[InputT, RuntimeContext], "NodeReturn[OutputT] | OutputT"
]
PlainNodeCallable = Callable[[InputT], "NodeReturn[OutputT] | OutputT"]
NodeCallable = (
Callable[[InputT, RuntimeContext], "NodeReturn[OutputT] | OutputT"]
| Callable[[InputT], "NodeReturn[OutputT] | OutputT"]
)
AsyncContextNodeCallable = Callable[ class ContextNodeCallable(Protocol[InputT, OutputT]):
[InputT, RuntimeContext], Awaitable["NodeReturn[OutputT] | OutputT"] def __call__(
self,
payload: InputT,
/,
ctx: RuntimeContext,
) -> "NodeReturn[OutputT] | OutputT": ...
class PlainNodeCallable(Protocol[InputT, OutputT]):
def __call__(self, payload: InputT, /) -> "NodeReturn[OutputT] | OutputT": ...
NodeCallable = ContextNodeCallable[InputT, OutputT] | PlainNodeCallable[
InputT, OutputT
] ]
AsyncPlainNodeCallable = Callable[
[InputT], Awaitable["NodeReturn[OutputT] | OutputT"]
] class AsyncContextNodeCallable(Protocol[InputT, OutputT]):
AsyncNodeCallable = ( def __call__(
Callable[ self,
[InputT, RuntimeContext], payload: InputT,
Awaitable["NodeReturn[OutputT] | OutputT"], /,
] ctx: RuntimeContext,
| Callable[[InputT], Awaitable["NodeReturn[OutputT] | OutputT"]] ) -> Awaitable["NodeReturn[OutputT] | OutputT"]: ...
)
class AsyncPlainNodeCallable(Protocol[InputT, OutputT]):
def __call__(
self,
payload: InputT,
/,
) -> Awaitable["NodeReturn[OutputT] | OutputT"]: ...
AsyncNodeCallable = AsyncContextNodeCallable[
InputT, OutputT
] | AsyncPlainNodeCallable[InputT, OutputT]
SyncRegistryHandler = Callable[[dict[str, Any], RuntimeContext], dict[str, Any]] SyncRegistryHandler = Callable[[dict[str, Any], RuntimeContext], dict[str, Any]]
AsyncRegistryHandler = Callable[ AsyncRegistryHandler = Callable[
@@ -166,8 +184,8 @@ class NodeSpec(Generic[InputT, OutputT]):
if self.accepts_context: if self.accepts_context:
if ctx is None: if ctx is None:
raise TypeError(f"node {self.name!r} requires RuntimeContext") raise TypeError(f"node {self.name!r} requires RuntimeContext")
return cast("ContextNodeCallable[InputT, OutputT]", self.fn)(payload, ctx) return cast(ContextNodeCallable[InputT, OutputT], self.fn)(payload, ctx)
return cast("PlainNodeCallable[InputT, OutputT]", self.fn)(payload) return cast(PlainNodeCallable[InputT, OutputT], self.fn)(payload)
def to_node_def(self) -> NodeDef: def to_node_def(self) -> NodeDef:
return NodeDef( return NodeDef(