misc changes: returns None, root as ., add

This commit is contained in:
lda
2026-05-18 00:45:08 +07:00 Verified
parent 6ee1d8f0cf
commit acd2e2ecb4
19 changed files with 274 additions and 53 deletions
+2 -1
View File
@@ -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",
+4 -4
View File
@@ -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 -1
View File
@@ -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
+6 -1
View File
@@ -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)
+4
View File
@@ -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]: ...
+9 -2
View File
@@ -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")