convert Callables to real protocols
ts the stuff langgraph do
This commit is contained in:
+40
-22
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user