misc changes: returns None, root as ., add
This commit is contained in:
@@ -13,7 +13,7 @@ from .callables import (
|
||||
from .inference import accepts_context, infer_models, is_basemodel_subclass
|
||||
from .decorator import node
|
||||
from .registry import build_async_registry, build_registry
|
||||
from .result import NodeReturn, Nothing, outcome
|
||||
from .result import NoOutput, NodeReturn, Nothing, outcome
|
||||
from .schema import schema_ref_for
|
||||
from .spec import NodeSpec
|
||||
|
||||
@@ -25,6 +25,7 @@ __all__ = [
|
||||
"ContextNodeCallable",
|
||||
"InputT",
|
||||
"NodeCallable",
|
||||
"NoOutput",
|
||||
"NodeReturn",
|
||||
"NodeSpec",
|
||||
"Nothing",
|
||||
|
||||
@@ -22,7 +22,7 @@ class ContextNodeCallable(Protocol[InputT_contra, OutputT_co]):
|
||||
payload: InputT_contra,
|
||||
/,
|
||||
ctx: RuntimeContext,
|
||||
) -> NodeReturn[OutputT_co] | OutputT_co: ...
|
||||
) -> NodeReturn[OutputT_co] | OutputT_co | None: ...
|
||||
|
||||
|
||||
class PlainNodeCallable(Protocol[InputT_contra, OutputT_co]):
|
||||
@@ -30,7 +30,7 @@ class PlainNodeCallable(Protocol[InputT_contra, OutputT_co]):
|
||||
self,
|
||||
payload: InputT_contra,
|
||||
/,
|
||||
) -> NodeReturn[OutputT_co] | OutputT_co: ...
|
||||
) -> NodeReturn[OutputT_co] | OutputT_co | None: ...
|
||||
|
||||
|
||||
NodeCallable = ContextNodeCallable[InputT, OutputT] | PlainNodeCallable[InputT, OutputT]
|
||||
@@ -42,7 +42,7 @@ class AsyncContextNodeCallable(Protocol[InputT_contra, OutputT_co]):
|
||||
payload: InputT_contra,
|
||||
/,
|
||||
ctx: RuntimeContext,
|
||||
) -> Awaitable[NodeReturn[OutputT_co] | OutputT_co]: ...
|
||||
) -> Awaitable[NodeReturn[OutputT_co] | OutputT_co | None]: ...
|
||||
|
||||
|
||||
class AsyncPlainNodeCallable(Protocol[InputT_contra, OutputT_co]):
|
||||
@@ -50,7 +50,7 @@ class AsyncPlainNodeCallable(Protocol[InputT_contra, OutputT_co]):
|
||||
self,
|
||||
payload: InputT_contra,
|
||||
/,
|
||||
) -> Awaitable[NodeReturn[OutputT_co] | OutputT_co]: ...
|
||||
) -> Awaitable[NodeReturn[OutputT_co] | OutputT_co | None]: ...
|
||||
|
||||
|
||||
AsyncNodeCallable = (
|
||||
|
||||
@@ -6,7 +6,12 @@ from typing import Any, cast, overload
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from .callables import AsyncNodeCallable, InputT, NodeCallable, OutputT
|
||||
from .callables import (
|
||||
AsyncNodeCallable,
|
||||
InputT,
|
||||
NodeCallable,
|
||||
OutputT,
|
||||
)
|
||||
from .inference import accepts_context, infer_models
|
||||
from .spec import NodeSpec
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ from pydantic import BaseModel
|
||||
|
||||
from wf_core import RuntimeContext
|
||||
|
||||
from .result import NodeReturn
|
||||
from .result import NodeReturn, Nothing
|
||||
|
||||
|
||||
def is_basemodel_subclass(value: object) -> bool:
|
||||
@@ -47,10 +47,15 @@ def infer_models(fn: Callable[..., object]) -> tuple[type[BaseModel], type[BaseM
|
||||
return_type = hints.get("return")
|
||||
if return_type is None:
|
||||
raise TypeError("node function must declare a return annotation")
|
||||
if return_type is type(None):
|
||||
return cast(type[BaseModel], input_model), Nothing
|
||||
|
||||
if is_basemodel_subclass(return_type):
|
||||
return cast(type[BaseModel], input_model), cast(type[BaseModel], return_type)
|
||||
|
||||
if return_type is NodeReturn:
|
||||
return cast(type[BaseModel], input_model), Nothing
|
||||
|
||||
origin = get_origin(return_type)
|
||||
if origin is NodeReturn:
|
||||
args = get_args(return_type)
|
||||
|
||||
@@ -20,6 +20,10 @@ class Nothing(BaseModel):
|
||||
"""Empty output model for nodes that only choose an outcome."""
|
||||
|
||||
|
||||
NoOutput = NodeReturn[Nothing]
|
||||
"""Type alias for outcome-only nodes that return no output payload."""
|
||||
|
||||
|
||||
@overload
|
||||
def outcome(name: str) -> NodeReturn[Nothing]: ...
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ from .callables import (
|
||||
PlainNodeCallable,
|
||||
SyncRegistryHandler,
|
||||
)
|
||||
from .result import NodeReturn
|
||||
from .result import NodeReturn, Nothing
|
||||
from .schema import schema_ref_for
|
||||
|
||||
|
||||
@@ -33,6 +33,8 @@ def _coerce_registry_result(
|
||||
default_outcome: str,
|
||||
raw: NodeReturn[BaseModel] | BaseModel,
|
||||
) -> dict[str, Any]:
|
||||
if raw is None and output_model is Nothing:
|
||||
return {"outcome": default_outcome, "output": {}}
|
||||
if isinstance(raw, NodeReturn):
|
||||
if not isinstance(raw.output, output_model):
|
||||
raise TypeError(
|
||||
@@ -69,7 +71,12 @@ class NodeSpec(Generic[InputT, OutputT]):
|
||||
self,
|
||||
payload: InputT,
|
||||
ctx: RuntimeContext | None = None,
|
||||
) -> NodeReturn[OutputT] | OutputT | Awaitable[NodeReturn[OutputT] | OutputT]:
|
||||
) -> (
|
||||
NodeReturn[OutputT]
|
||||
| OutputT
|
||||
| None
|
||||
| Awaitable[NodeReturn[OutputT] | OutputT | None]
|
||||
):
|
||||
if self.accepts_context:
|
||||
if ctx is None:
|
||||
raise TypeError(f"node {self.name!r} requires RuntimeContext")
|
||||
|
||||
Reference in New Issue
Block a user