From 720dfd3b095ff403843e911bb6af536af2dfe098 Mon Sep 17 00:00:00 2001 From: Christophe Bornet Date: Wed, 10 Jun 2026 22:16:03 +0200 Subject: [PATCH] chore(core): improve typing of Runnable `__or__` (#34530) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `Runnable.__or__`, `Runnable.__ror__`, and their `RunnableSequence` and `StructuredPrompt` overrides previously erased composition types: the right-hand operand was typed `Runnable[Any, Other]`, so piping two runnables together always produced `RunnableSerializable[Input, Any]`. Type information was lost at every `|`, which is why chains so often needed a `chain: Runnable = ...` annotation just to recover usable inference. This adds `@overload`s so the `Output` of one step flows into the `Input` of the next and the composed result carries the real `Output` type through. `Runnable[int, str] | Runnable[str, float]` now infers `RunnableSerializable[int, float]` instead of `[int, Any]`. `coerce_to_runnable` gains overloads so a `Mapping` resolves to `RunnableParallel` while everything else stays a `Runnable`. As a knock-on effect, dozens of now-unnecessary `: Runnable` annotations were dropped from the test suite. Runtime behavior is unchanged — this is a typing-only change. ## Impact on type-checked code Most users will simply get better inference. Two changes can require a small adjustment if you run a type checker (`mypy`, `pyright`): ### Stricter operand matching in `|` The right-hand side of `|` is now typed `Runnable[Output, Other]` rather than `Runnable[Any, Other]`, so the right operand's declared **input** must match the left operand's **output**. This is more accurate, but it surfaces a common pattern that was previously silent: piping a step that outputs a plain `dict` into a step whose declared input is a more specific type (for example a `TypedDict`). It still works at runtime; the checker now reports an `[operator]` error. If you hit this, narrow the boundary with a `cast` (or an explicit annotation): ```python from typing import Any, cast from langchain_core.runnables import Runnable # upstream outputs a dict; downstream declares a narrower input type chain = cast("Runnable[Any, MyInput]", upstream) | downstream ``` ### `list` → `Sequence` on `RunnableEach` / `map()` `Runnable.map()` and the `invoke` / `ainvoke` methods of `RunnableEach` now accept `Sequence[Input]` instead of `list[Input]`. Callers are unaffected — a `list` is a `Sequence`, and tuples or other sequences now type-check too. The only thing to adjust: if you **subclass** `RunnableEach` (or `RunnableEachBase`) and override these methods with a `list[...]` parameter, widen the annotation to `Sequence[...]` so the override stays compatible with the base signature. --------- Co-authored-by: Mason Daugherty --- .../core/langchain_core/prompts/structured.py | 44 ++++- libs/core/langchain_core/runnables/base.py | 180 ++++++++++++++---- libs/core/langchain_core/runnables/branch.py | 4 +- .../tests/unit_tests/runnables/test_graph.py | 2 +- .../unit_tests/runnables/test_runnable.py | 136 +++++++------ .../runnables/test_runnable_events_v2.py | 2 +- .../runnables/openai_functions.py | 13 +- .../tests/unit_tests/agents/test_agent.py | 6 +- 8 files changed, 273 insertions(+), 114 deletions(-) diff --git a/libs/core/langchain_core/prompts/structured.py b/libs/core/langchain_core/prompts/structured.py index 0150007984..77aeef00e1 100644 --- a/libs/core/langchain_core/prompts/structured.py +++ b/libs/core/langchain_core/prompts/structured.py @@ -1,8 +1,16 @@ """Structured prompt template for a language model.""" -from collections.abc import AsyncIterator, Callable, Iterator, Mapping, Sequence +from collections.abc import ( + AsyncIterator, + Awaitable, + Callable, + Iterator, + Mapping, + Sequence, +) from typing import ( Any, + overload, ) from pydantic import BaseModel, Field @@ -10,6 +18,7 @@ from typing_extensions import override from langchain_core._api.beta_decorator import beta from langchain_core.language_models.base import BaseLanguageModel +from langchain_core.prompt_values import PromptValue from langchain_core.prompts.chat import ( ChatPromptTemplate, MessageLikeRepresentation, @@ -135,15 +144,36 @@ class StructuredPrompt(ChatPromptTemplate): """ return cls(messages, schema, **kwargs) + @overload + def __or__( + self, other: Mapping[str, Any] + ) -> RunnableSerializable[dict[str, Any], dict[str, Any]]: ... + + @overload + def __or__( + self, + other: Callable[[PromptValue], Runnable[PromptValue, Other]] + | Callable[[PromptValue], Awaitable[Runnable[PromptValue, Other]]], + ) -> RunnableSerializable[dict[str, Any], Other]: ... + + @overload + def __or__( + self, + other: Runnable[PromptValue, Other] + | Callable[[Iterator[PromptValue]], Iterator[Other]] + | Callable[[AsyncIterator[PromptValue]], AsyncIterator[Other]] + | Callable[[PromptValue], Other], + ) -> RunnableSerializable[dict[str, Any], Other]: ... + @override def __or__( self, - other: Runnable[Any, Other] - | Callable[[Iterator[Any]], Iterator[Other]] - | Callable[[AsyncIterator[Any]], AsyncIterator[Other]] - | Callable[[Any], Other] - | Mapping[str, Runnable[Any, Other] | Callable[[Any], Other] | Any], - ) -> RunnableSerializable[dict[str, Any], Other]: + other: Runnable[PromptValue, Other] + | Callable[[Iterator[PromptValue]], Iterator[Other]] + | Callable[[AsyncIterator[PromptValue]], AsyncIterator[Other]] + | Callable[[PromptValue], Other] + | Mapping[str, Runnable[PromptValue, Any] | Callable[[PromptValue], Any] | Any], + ) -> RunnableSerializable[dict[str, Any], Any]: return self.pipe(other) def pipe( diff --git a/libs/core/langchain_core/runnables/base.py b/libs/core/langchain_core/runnables/base.py index 65a53fb1d6..f6b47558af 100644 --- a/libs/core/langchain_core/runnables/base.py +++ b/libs/core/langchain_core/runnables/base.py @@ -616,14 +616,35 @@ class Runnable(ABC, Generic[Input, Output]): if isinstance(node.data, BasePromptTemplate) ] + @overload + def __or__( + self, other: Mapping[str, Any] + ) -> RunnableSerializable[Input, dict[str, Any]]: ... + + @overload def __or__( self, - other: Runnable[Any, Other] - | Callable[[Iterator[Any]], Iterator[Other]] - | Callable[[AsyncIterator[Any]], AsyncIterator[Other]] - | Callable[[Any], Other] - | Mapping[str, Runnable[Any, Other] | Callable[[Any], Other] | Any], - ) -> RunnableSerializable[Input, Other]: + other: Callable[[Output], Runnable[Output, Other]] + | Callable[[Output], Awaitable[Runnable[Output, Other]]], + ) -> RunnableSerializable[Input, Other]: ... + + @overload + def __or__( + self, + other: Runnable[Output, Other] + | Callable[[Iterator[Output]], Iterator[Other]] + | Callable[[AsyncIterator[Output]], AsyncIterator[Other]] + | Callable[[Output], Other], + ) -> RunnableSerializable[Input, Other]: ... + + def __or__( + self, + other: Runnable[Output, Other] + | Callable[[Iterator[Output]], Iterator[Other]] + | Callable[[AsyncIterator[Output]], AsyncIterator[Other]] + | Callable[[Output], Other] + | Mapping[str, Runnable[Output, Any] | Callable[[Output], Any] | Any], + ) -> RunnableSerializable[Input, Any]: """Runnable "or" operator. Compose this `Runnable` with another object to create a @@ -637,14 +658,36 @@ class Runnable(ABC, Generic[Input, Output]): """ return RunnableSequence(self, coerce_to_runnable(other)) + @overload def __ror__( self, - other: Runnable[Other, Any] - | Callable[[Iterator[Other]], Iterator[Any]] - | Callable[[AsyncIterator[Other]], AsyncIterator[Any]] + other: Mapping[str, Any], + ) -> RunnableSerializable[Any, Output]: ... + + @overload + def __ror__( + self, + other: Callable[[Other], Runnable[Other, Input]] + | Callable[[Other], Awaitable[Runnable[Other, Input]]], + ) -> RunnableSerializable[Other, Output]: ... + + @overload + def __ror__( + self, + other: Runnable[Other, Input] + | Callable[[Iterator[Other]], Iterator[Input]] + | Callable[[AsyncIterator[Other]], AsyncIterator[Input]] + | Callable[[Other], Input], + ) -> RunnableSerializable[Other, Output]: ... + + def __ror__( + self, + other: Runnable[Other, Input] + | Callable[[Iterator[Other]], Iterator[Input]] + | Callable[[AsyncIterator[Other]], AsyncIterator[Input]] | Callable[[Other], Any] - | Mapping[str, Runnable[Other, Any] | Callable[[Other], Any] | Any], - ) -> RunnableSerializable[Other, Output]: + | Mapping[str, Runnable[Other, Input] | Callable[[Other], Any] | Any], + ) -> RunnableSerializable[Any, Output]: """Runnable "reverse-or" operator. Compose this `Runnable` with another object to create a @@ -768,7 +811,7 @@ class Runnable(ABC, Generic[Input, Output]): # Import locally to prevent circular import from langchain_core.runnables.passthrough import RunnablePick # noqa: PLC0415 - return self | RunnablePick(keys) + return self | RunnablePick(keys) # type: ignore[operator, no-any-return] def assign( self, @@ -815,7 +858,7 @@ class Runnable(ABC, Generic[Input, Output]): # Import locally to prevent circular import from langchain_core.runnables.passthrough import RunnableAssign # noqa: PLC0415 - return self | RunnableAssign(RunnableParallel[dict[str, Any]](kwargs)) + return self | RunnableAssign(RunnableParallel[dict[str, Any]](kwargs)) # type: ignore[operator, no-any-return] """ --- Public API --- """ @@ -2099,7 +2142,7 @@ class Runnable(ABC, Generic[Input, Output]): exponential_jitter_params=exponential_jitter_params, ) - def map(self) -> Runnable[list[Input], list[Output]]: + def map(self) -> Runnable[Sequence[Input], list[Output]]: """Return a new `Runnable` that maps a list of inputs to a list of outputs. Calls `invoke` with each input. @@ -3251,15 +3294,36 @@ class RunnableSequence(RunnableSerializable[Input, Output]): for i, s in enumerate(self.steps) ) + @overload + def __or__( + self, other: Mapping[str, Any] + ) -> RunnableSerializable[Input, dict[str, Any]]: ... + + @overload + def __or__( + self, + other: Callable[[Output], Runnable[Output, Other]] + | Callable[[Output], Awaitable[Runnable[Output, Other]]], + ) -> RunnableSerializable[Input, Other]: ... + + @overload + def __or__( + self, + other: Runnable[Output, Other] + | Callable[[Iterator[Output]], Iterator[Other]] + | Callable[[AsyncIterator[Output]], AsyncIterator[Other]] + | Callable[[Output], Other], + ) -> RunnableSerializable[Input, Other]: ... + @override def __or__( self, - other: Runnable[Any, Other] - | Callable[[Iterator[Any]], Iterator[Other]] - | Callable[[AsyncIterator[Any]], AsyncIterator[Other]] - | Callable[[Any], Other] - | Mapping[str, Runnable[Any, Other] | Callable[[Any], Other] | Any], - ) -> RunnableSerializable[Input, Other]: + other: Runnable[Output, Other] + | Callable[[Iterator[Output]], Iterator[Other]] + | Callable[[AsyncIterator[Output]], AsyncIterator[Other]] + | Callable[[Output], Other] + | Mapping[str, Runnable[Output, Any] | Callable[[Output], Any] | Any], + ) -> RunnableSerializable[Input, Any]: if isinstance(other, RunnableSequence): return RunnableSequence( self.first, @@ -3278,14 +3342,36 @@ class RunnableSequence(RunnableSerializable[Input, Output]): name=self.name, ) + @overload + def __ror__( + self, + other: Mapping[str, Any], + ) -> RunnableSerializable[Any, Output]: ... + + @overload + def __ror__( + self, + other: Callable[[Other], Runnable[Other, Input]] + | Callable[[Other], Awaitable[Runnable[Other, Input]]], + ) -> RunnableSerializable[Other, Output]: ... + + @overload + def __ror__( + self, + other: Runnable[Other, Input] + | Callable[[Iterator[Other]], Iterator[Input]] + | Callable[[AsyncIterator[Other]], AsyncIterator[Input]] + | Callable[[Other], Input], + ) -> RunnableSerializable[Other, Output]: ... + @override def __ror__( self, - other: Runnable[Other, Any] - | Callable[[Iterator[Other]], Iterator[Any]] - | Callable[[AsyncIterator[Other]], AsyncIterator[Any]] + other: Runnable[Other, Input] + | Callable[[Iterator[Other]], Iterator[Input]] + | Callable[[AsyncIterator[Other]], AsyncIterator[Input]] | Callable[[Other], Any] - | Mapping[str, Runnable[Other, Any] | Callable[[Other], Any] | Any], + | Mapping[str, Runnable[Other, Input] | Callable[[Other], Any] | Any], ) -> RunnableSerializable[Other, Output]: if isinstance(other, RunnableSequence): return RunnableSequence( @@ -5449,7 +5535,7 @@ class RunnableLambda(Runnable[Input, Output]): yield chunk -class RunnableEachBase(RunnableSerializable[list[Input], list[Output]]): +class RunnableEachBase(RunnableSerializable[Sequence[Input], list[Output]]): """RunnableEachBase class. `Runnable` that calls another `Runnable` for each element of the input sequence. @@ -5540,7 +5626,7 @@ class RunnableEachBase(RunnableSerializable[list[Input], list[Output]]): def _invoke( self, - inputs: list[Input], + inputs: Sequence[Input], run_manager: CallbackManagerForChainRun, config: RunnableConfig, **kwargs: Any, @@ -5548,17 +5634,20 @@ class RunnableEachBase(RunnableSerializable[list[Input], list[Output]]): configs = [ patch_config(config, callbacks=run_manager.get_child()) for _ in inputs ] - return self.bound.batch(inputs, configs, **kwargs) + return self.bound.batch(list(inputs), configs, **kwargs) @override def invoke( - self, input: list[Input], config: RunnableConfig | None = None, **kwargs: Any + self, + input: Sequence[Input], + config: RunnableConfig | None = None, + **kwargs: Any, ) -> list[Output]: return self._call_with_config(self._invoke, input, config, **kwargs) async def _ainvoke( self, - inputs: list[Input], + inputs: Sequence[Input], run_manager: AsyncCallbackManagerForChainRun, config: RunnableConfig, **kwargs: Any, @@ -5566,13 +5655,18 @@ class RunnableEachBase(RunnableSerializable[list[Input], list[Output]]): configs = [ patch_config(config, callbacks=run_manager.get_child()) for _ in inputs ] - return await self.bound.abatch(inputs, configs, **kwargs) + return await self.bound.abatch(list(inputs), configs, **kwargs) @override async def ainvoke( - self, input: list[Input], config: RunnableConfig | None = None, **kwargs: Any + self, + input: Sequence[Input], + config: RunnableConfig | None = None, + **kwargs: Any, ) -> list[Output]: - return await self._acall_with_config(self._ainvoke, input, config, **kwargs) + return await self._acall_with_config( + self._ainvoke, input, config=config, **kwargs + ) @override def astream_events( # type: ignore[override] @@ -6488,7 +6582,25 @@ RunnableLike = ( ) -def coerce_to_runnable(thing: RunnableLike[Input, Output]) -> Runnable[Input, Output]: +@overload +def coerce_to_runnable( + thing: Runnable[Input, Output] + | Callable[[Input], Output] + | Callable[[Input], Awaitable[Output]] + | Callable[[Iterator[Input]], Iterator[Output]] + | Callable[[AsyncIterator[Input]], AsyncIterator[Output]] + | _RunnableCallableSync[Input, Output] + | _RunnableCallableAsync[Input, Output] + | _RunnableCallableIterator[Input, Output] + | _RunnableCallableAsyncIterator[Input, Output], +) -> Runnable[Input, Output]: ... + + +@overload +def coerce_to_runnable(thing: Mapping[str, Any]) -> RunnableParallel[Input]: ... + + +def coerce_to_runnable(thing: RunnableLike[Input, Output]) -> Runnable[Input, Any]: """Coerce a `Runnable`-like object into a `Runnable`. Args: @@ -6507,7 +6619,7 @@ def coerce_to_runnable(thing: RunnableLike[Input, Output]) -> Runnable[Input, Ou if callable(thing): return RunnableLambda(cast("Callable[[Input], Output]", thing)) if isinstance(thing, dict): - return cast("Runnable[Input, Output]", RunnableParallel(thing)) + return RunnableParallel(thing) msg = ( f"Expected a Runnable, callable or dict." f"Instead got an unsupported type: {type(thing)}" diff --git a/libs/core/langchain_core/runnables/branch.py b/libs/core/langchain_core/runnables/branch.py index b8a3ffb7d1..277a4a3607 100644 --- a/libs/core/langchain_core/runnables/branch.py +++ b/libs/core/langchain_core/runnables/branch.py @@ -106,7 +106,7 @@ class RunnableBranch(RunnableSerializable[Input, Output]): msg = "RunnableBranch default must be Runnable, callable or mapping." raise TypeError(msg) - default_ = coerce_to_runnable(cast("RunnableLike[Input, Output]", default)) + default_ = coerce_to_runnable(cast("Runnable[Input, Output]", default)) branches_ = [] @@ -126,7 +126,7 @@ class RunnableBranch(RunnableSerializable[Input, Output]): raise ValueError(msg) condition, runnable = branch condition = cast("Runnable[Input, bool]", coerce_to_runnable(condition)) - runnable = coerce_to_runnable(runnable) + runnable = coerce_to_runnable(cast("Runnable[Input, Output]", runnable)) branches_.append((condition, runnable)) super().__init__( diff --git a/libs/core/tests/unit_tests/runnables/test_graph.py b/libs/core/tests/unit_tests/runnables/test_graph.py index ffd90a0779..df2a54f59e 100644 --- a/libs/core/tests/unit_tests/runnables/test_graph.py +++ b/libs/core/tests/unit_tests/runnables/test_graph.py @@ -233,7 +233,7 @@ def test_graph_sequence_map(snapshot: SnapshotAssertion) -> None: return str_parser return xml_parser - sequence: Runnable = ( + sequence = ( prompt | fake_llm | { diff --git a/libs/core/tests/unit_tests/runnables/test_runnable.py b/libs/core/tests/unit_tests/runnables/test_runnable.py index 1cccc2743b..25da4d3b15 100644 --- a/libs/core/tests/unit_tests/runnables/test_runnable.py +++ b/libs/core/tests/unit_tests/runnables/test_runnable.py @@ -461,7 +461,7 @@ def test_schemas(snapshot: SnapshotAssertion) -> None: } assert router.get_output_jsonschema() == {"title": "RouterRunnableOutput"} - seq_w_map: Runnable = ( + seq_w_map = ( prompt | fake_llm | { @@ -530,7 +530,7 @@ def test_passthrough_assign_schema() -> None: invalid_seq_w_assign = ( RunnablePassthrough.assign(context=itemgetter("question") | retriever) - | fake_llm + | fake_llm # type: ignore[operator] ) # fallback to RunnableAssign.input_schema if next runnable doesn't have @@ -673,7 +673,7 @@ def test_schema_with_itemgetter() -> None: "type": "object", } prompt = ChatPromptTemplate.from_template("what is {language}?") - chain: Runnable = {"language": itemgetter("language")} | prompt + chain = {"language": itemgetter("language")} | prompt assert _schema(chain.input_schema) == { "properties": {"language": {"title": "Language"}}, "required": ["language"], @@ -696,7 +696,7 @@ def test_schema_complex_seq() -> None: assert chain1.name == "city_chain" - chain2: Runnable = ( + chain2 = ( {"city": chain1, "language": itemgetter("language")} | prompt2 | model @@ -840,7 +840,7 @@ def test_configurable_fields(snapshot: SnapshotAssertion) -> None: "required": ["lang", "name"], } - chain_with_map_configurable: Runnable = prompt_configurable | { + chain_with_map_configurable = prompt_configurable | { "llm1": fake_llm_configurable | StrOutputParser(), "llm2": fake_llm_configurable | StrOutputParser(), "llm3": fake_llm.configurable_fields( @@ -1708,7 +1708,7 @@ def test_with_listener_propagation(mocker: MockerFixture) -> None: + "{question}" ) chat = FakeListChatModel(responses=["foo"]) - chain: Runnable = prompt | chat + chain = prompt | chat mock_start = mocker.Mock() mock_end = mocker.Mock() chain_with_listeners = chain.with_listeners(on_start=mock_start, on_end=mock_end) @@ -1722,9 +1722,7 @@ def test_with_listener_propagation(mocker: MockerFixture) -> None: mock_start.reset_mock() mock_end.reset_mock() - chain_with_listeners.with_types(output_type=str).invoke( - {"question": "Who are you?"} - ) + chain_with_listeners.invoke({"question": "Who are you?"}) assert mock_start.call_count == 1 assert mock_start.call_args[0][0].name == "RunnableSequence" @@ -2473,7 +2471,7 @@ async def test_stream_log_retriever() -> None: ) llm = FakeListLLM(responses=["foo", "bar"]) - chain: Runnable = ( + chain = ( {"documents": FakeRetriever(), "question": itemgetter("question")} | prompt | {"one": llm, "two": llm} @@ -2740,7 +2738,7 @@ Question: parser = CommaSeparatedListOutputParser() - chain: Runnable = ( + chain = ( { "question": RunnablePassthrough[str]() | passthrough, "documents": passthrough | retriever, @@ -2870,7 +2868,7 @@ def test_router_runnable(mocker: MockerFixture, snapshot: SnapshotAssertion) -> "You are an english major. Answer the question: {question}" ) | FakeListLLM(responses=["2"]) router = RouterRunnable({"math": chain1, "english": chain2}) - chain: Runnable = { + chain = { "key": lambda x: x["key"], "input": {"question": lambda x: x["question"]}, } | router @@ -2914,7 +2912,7 @@ async def test_router_runnable_async() -> None: "You are an english major. Answer the question: {question}" ) | FakeListLLM(responses=["2"]) router = RouterRunnable({"math": chain1, "english": chain2}) - chain: Runnable = { + chain = { "key": lambda x: x["key"], "input": {"question": lambda x: x["question"]}, } | router @@ -2946,7 +2944,7 @@ def test_higher_order_lambda_runnable( input={"question": lambda x: x["question"]}, ) - def router(params: dict[str, Any]) -> Runnable: + def router(params: dict[str, Any]) -> Runnable[dict[str, Any], str]: if params["key"] == "math": return itemgetter("input") | math_chain if params["key"] == "english": @@ -2954,7 +2952,7 @@ def test_higher_order_lambda_runnable( msg = f"Unknown key: {params['key']}" raise ValueError(msg) - chain: Runnable = input_map | router + chain = input_map | router assert dumps(chain, pretty=True) == snapshot result = chain.invoke({"key": "math", "question": "2 + 2"}) @@ -3010,7 +3008,7 @@ async def test_higher_order_lambda_runnable_async(mocker: MockerFixture) -> None msg = f"Unknown key: {value['key']}" raise ValueError(msg) - chain: Runnable = input_map | router + chain = input_map | router result = await chain.ainvoke({"key": "math", "question": "2 + 2"}) assert result == "4" @@ -3032,7 +3030,7 @@ async def test_higher_order_lambda_runnable_async(mocker: MockerFixture) -> None msg = f"Unknown key: {params['key']}" raise ValueError(msg) - achain: Runnable = input_map | arouter + achain = input_map | arouter math_spy = mocker.spy(math_chain.__class__, "ainvoke") tracer = FakeTracer() assert ( @@ -3139,7 +3137,7 @@ def test_map_stream() -> None: # sleep to better simulate a real stream llm = FakeStreamingListLLM(responses=[llm_res], sleep=0.01) - chain: Runnable = prompt | { + chain = prompt | { "chat": chat.bind(stop=["Thought:"]), "llm": llm, "passthrough": RunnablePassthrough(), @@ -3150,6 +3148,7 @@ def test_map_stream() -> None: final_value = None streamed_chunks = [] for chunk in stream: + assert isinstance(chunk, AddableDict) streamed_chunks.append(chunk) if final_value is None: final_value = chunk @@ -3164,7 +3163,9 @@ def test_map_stream() -> None: assert len(streamed_chunks) == len(chat_res) + len(llm_res) + 1 assert all(len(c.keys()) == 1 for c in streamed_chunks) assert final_value is not None - assert final_value.get("chat").content == "i'm a chatbot" + chat_message = final_value.get("chat") + assert chat_message is not None + assert chat_message.content == "i'm a chatbot" assert final_value.get("llm") == "i'm a textbot" assert final_value.get("passthrough") == prompt.invoke( {"question": "What is your name?"} @@ -3177,19 +3178,19 @@ def test_map_stream() -> None: "type": "string", } - stream = chain_pick_one.stream({"question": "What is your name?"}) + stream_picked = chain_pick_one.stream({"question": "What is your name?"}) final_value = None - streamed_chunks = [] - for chunk in stream: - streamed_chunks.append(chunk) + streamed_chunks_picked = [] + for chunk in stream_picked: + streamed_chunks_picked.append(chunk) if final_value is None: final_value = chunk else: final_value += chunk - assert streamed_chunks[0] == "i" - assert len(streamed_chunks) == len(llm_res) + assert streamed_chunks_picked[0] == "i" + assert len(streamed_chunks_picked) == len(llm_res) chain_pick_two = chain.assign(hello=RunnablePick("llm").pipe(llm)).pick( [ @@ -3208,30 +3209,30 @@ def test_map_stream() -> None: "required": ["llm", "hello"], } - stream = chain_pick_two.stream({"question": "What is your name?"}) + stream_picked = chain_pick_two.stream({"question": "What is your name?"}) final_value = None - streamed_chunks = [] - for chunk in stream: - streamed_chunks.append(chunk) + streamed_chunks_picked = [] + for chunk in stream_picked: + streamed_chunks_picked.append(chunk) if final_value is None: final_value = chunk else: final_value += chunk - assert streamed_chunks[0] in [ + assert streamed_chunks_picked[0] in [ {"llm": "i"}, {"chat": _any_id_ai_message_chunk(content="i")}, ] if not ( # TODO: Rewrite properly the statement above - streamed_chunks[0] == {"llm": "i"} + streamed_chunks_picked[0] == {"llm": "i"} or {"chat": _any_id_ai_message_chunk(content="i")} ): - msg = f"Got an unexpected chunk: {streamed_chunks[0]}" + msg = f"Got an unexpected chunk: {streamed_chunks_picked[0]}" raise AssertionError(msg) - assert len(streamed_chunks) == len(llm_res) + len(chat_res) + assert len(streamed_chunks_picked) == len(llm_res) + len(chat_res) def test_map_stream_iterator_input() -> None: @@ -3248,7 +3249,7 @@ def test_map_stream_iterator_input() -> None: # sleep to better simulate a real stream llm = FakeStreamingListLLM(responses=[llm_res], sleep=0.01) - chain: Runnable = ( + chain = ( prompt | llm | { @@ -3263,6 +3264,7 @@ def test_map_stream_iterator_input() -> None: final_value = None streamed_chunks = [] for chunk in stream: + assert isinstance(chunk, AddableDict) streamed_chunks.append(chunk) if final_value is None: final_value = chunk @@ -3277,7 +3279,9 @@ def test_map_stream_iterator_input() -> None: assert len(streamed_chunks) == len(chat_res) + len(llm_res) + len(llm_res) assert all(len(c.keys()) == 1 for c in streamed_chunks) assert final_value is not None - assert final_value.get("chat").content == "i'm a chatbot" + chat_message = final_value.get("chat") + assert chat_message is not None + assert chat_message.content == "i'm a chatbot" assert final_value.get("llm") == "i'm a textbot" assert final_value.get("passthrough") == "i'm a textbot" @@ -3296,7 +3300,7 @@ async def test_map_astream() -> None: # sleep to better simulate a real stream llm = FakeStreamingListLLM(responses=[llm_res], sleep=0.01) - chain: Runnable = prompt | { + chain = prompt | { "chat": chat.bind(stop=["Thought:"]), "llm": llm, "passthrough": RunnablePassthrough(), @@ -3307,6 +3311,7 @@ async def test_map_astream() -> None: final_value = None streamed_chunks = [] async for chunk in stream: + assert isinstance(chunk, AddableDict) streamed_chunks.append(chunk) if final_value is None: final_value = chunk @@ -3321,8 +3326,10 @@ async def test_map_astream() -> None: assert len(streamed_chunks) == len(chat_res) + len(llm_res) + 1 assert all(len(c.keys()) == 1 for c in streamed_chunks) assert final_value is not None - assert final_value.get("chat").content == "i'm a chatbot" - final_value["chat"].id = AnyStr() + chat_message = final_value.get("chat") + assert chat_message is not None + assert chat_message.content == "i'm a chatbot" + chat_message.id = AnyStr() assert final_value.get("llm") == "i'm a textbot" assert final_value.get("passthrough") == prompt.invoke( {"question": "What is your name?"} @@ -3332,12 +3339,12 @@ async def test_map_astream() -> None: final_state = None streamed_ops = [] - async for chunk in chain.astream_log({"question": "What is your name?"}): - streamed_ops.extend(chunk.ops) + async for patch in chain.astream_log({"question": "What is your name?"}): + streamed_ops.extend(patch.ops) if final_state is None: - final_state = chunk + final_state = patch else: - final_state += chunk + final_state += patch final_state = cast("RunLog", final_state) assert final_state.state["final_output"] == final_value @@ -3365,13 +3372,13 @@ async def test_map_astream() -> None: # Test astream_log with include filters final_state = None - async for chunk in chain.astream_log( + async for patch in chain.astream_log( {"question": "What is your name?"}, include_names=["FakeListChatModel"] ): if final_state is None: - final_state = chunk + final_state = patch else: - final_state += chunk + final_state += patch final_state = cast("RunLog", final_state) assert final_state.state["final_output"] == final_value @@ -3381,13 +3388,13 @@ async def test_map_astream() -> None: # Test astream_log with exclude filters final_state = None - async for chunk in chain.astream_log( + async for patch in chain.astream_log( {"question": "What is your name?"}, exclude_names=["FakeListChatModel"] ): if final_state is None: - final_state = chunk + final_state = patch else: - final_state += chunk + final_state += patch final_state = cast("RunLog", final_state) assert final_state.state["final_output"] == final_value @@ -3425,7 +3432,7 @@ async def test_map_astream_iterator_input() -> None: # sleep to better simulate a real stream llm = FakeStreamingListLLM(responses=[llm_res], sleep=0.01) - chain: Runnable = ( + chain = ( prompt | llm | { @@ -3440,6 +3447,7 @@ async def test_map_astream_iterator_input() -> None: final_value = None streamed_chunks = [] async for chunk in stream: + assert isinstance(chunk, AddableDict) streamed_chunks.append(chunk) if final_value is None: final_value = chunk @@ -3454,7 +3462,9 @@ async def test_map_astream_iterator_input() -> None: assert len(streamed_chunks) == len(chat_res) + len(llm_res) + len(llm_res) assert all(len(c.keys()) == 1 for c in streamed_chunks) assert final_value is not None - assert final_value.get("chat").content == "i'm a chatbot" + chat_message = final_value.get("chat") + assert chat_message is not None + assert chat_message.content == "i'm a chatbot" assert final_value.get("llm") == "i'm a textbot" assert final_value.get("passthrough") == llm_res @@ -3555,11 +3565,11 @@ def test_deep_stream_assign() -> None: ) llm = FakeStreamingListLLM(responses=["foo-lish"]) - chain: Runnable = prompt | llm | {"str": StrOutputParser()} + chain = prompt | llm | {"str": StrOutputParser()} stream = chain.stream({"question": "What up"}) - chunks = list(stream) + chunks = [chunk for chunk in stream if isinstance(chunk, AddableDict)] assert len(chunks) == len("foo-lish") assert add(chunks) == {"str": "foo-lish"} @@ -3677,14 +3687,16 @@ async def test_deep_astream_assign() -> None: ) llm = FakeStreamingListLLM(responses=["foo-lish"]) - chain: Runnable = prompt | llm | {"str": StrOutputParser()} + chain = prompt | llm | {"str": StrOutputParser()} stream = chain.astream({"question": "What up"}) - chunks = [chunk async for chunk in stream] + chunks: list[AddableDict] = [ + chunk async for chunk in stream if isinstance(chunk, AddableDict) + ] assert len(chunks) == len("foo-lish") - assert add(chunks) == {"str": "foo-lish"} + assert add(chunks) == AddableDict({"str": "foo-lish"}) chain_with_assign = chain.assign( hello=itemgetter("str") | llm, @@ -3758,9 +3770,11 @@ async def test_deep_astream_assign() -> None: "required": ["str", "hello"], } - chunks = [] - async for chunk in chain_with_assign_shadow.astream({"question": "What up"}): - chunks.append(chunk) + chunks = [ + chunk + async for chunk in chain_with_assign_shadow.astream({"question": "What up"}) + if isinstance(chunk, AddableDict) + ] assert len(chunks) == len("foo-lish") + 1 assert add(chunks) == {"str": "shadow", "hello": "foo-lish"} @@ -3860,6 +3874,8 @@ def test_each_simple() -> None: [["a", "b"], ["c"]], [["c", "e"]], ] + # `.map()` accepts any `Sequence`, not just `list` (e.g. a tuple). + assert parser.map().invoke(("a, b", "c")) == [["a", "b"], ["c"]] def test_each(snapshot: SnapshotAssertion) -> None: @@ -5354,8 +5370,8 @@ async def test_runnable_gen_transform() -> None: async for i in ints: yield i + 1 - chain: Runnable = RunnableGenerator(gen_indexes, agen_indexes) | plus_one - achain: Runnable = RunnableGenerator(gen_indexes, agen_indexes) | aplus_one + chain = RunnableGenerator(gen_indexes, agen_indexes) | plus_one + achain = RunnableGenerator(gen_indexes, agen_indexes) | aplus_one assert chain.get_input_jsonschema() == { "title": "gen_indexes_input", diff --git a/libs/core/tests/unit_tests/runnables/test_runnable_events_v2.py b/libs/core/tests/unit_tests/runnables/test_runnable_events_v2.py index f59500dba4..12573fd67d 100644 --- a/libs/core/tests/unit_tests/runnables/test_runnable_events_v2.py +++ b/libs/core/tests/unit_tests/runnables/test_runnable_events_v2.py @@ -93,7 +93,7 @@ async def _collect_events( async def test_event_stream_with_simple_function_tool() -> None: """Test the event stream with a function and tool.""" - def foo(x: int) -> dict[str, Any]: + def foo(x: int) -> dict[str, int]: """Foo.""" _ = x return {"x": 5} diff --git a/libs/langchain/langchain_classic/runnables/openai_functions.py b/libs/langchain/langchain_classic/runnables/openai_functions.py index 29f2194c85..b1671573aa 100644 --- a/libs/langchain/langchain_classic/runnables/openai_functions.py +++ b/libs/langchain/langchain_classic/runnables/openai_functions.py @@ -1,10 +1,10 @@ from collections.abc import Callable, Mapping from operator import itemgetter -from typing import Any +from typing import Any, cast from langchain_core.messages import BaseMessage from langchain_core.output_parsers.openai_functions import JsonOutputFunctionsParser -from langchain_core.runnables import RouterRunnable, Runnable +from langchain_core.runnables import RouterInput, RouterRunnable, Runnable from langchain_core.runnables.base import RunnableBindingBase from typing_extensions import TypedDict @@ -46,9 +46,10 @@ class OpenAIFunctionsRouter(RunnableBindingBase[BaseMessage, Any]): # type: ign if not all(func["name"] in runnables for func in functions): msg = "One or more function names are not found in runnables." raise ValueError(msg) - router = ( + router = cast( + "Runnable[Any, RouterInput]", JsonOutputFunctionsParser(args_only=False) - | {"key": itemgetter("name"), "input": itemgetter("arguments")} - | RouterRunnable(runnables) - ) + | {"key": itemgetter("name"), "input": itemgetter("arguments")}, + ) | RouterRunnable(runnables) + super().__init__(bound=router, kwargs={}, functions=functions) diff --git a/libs/langchain/tests/unit_tests/agents/test_agent.py b/libs/langchain/tests/unit_tests/agents/test_agent.py index 8cbcb19856..5317f0d180 100644 --- a/libs/langchain/tests/unit_tests/agents/test_agent.py +++ b/libs/langchain/tests/unit_tests/agents/test_agent.py @@ -463,7 +463,7 @@ async def test_runnable_agent() -> None: ], ) - def fake_parse(_: dict) -> AgentFinish | AgentAction: + def fake_parse(_: AIMessage) -> AgentFinish | AgentAction: """A parser.""" return AgentFinish(return_values={"foo": "meow"}, log="hard-coded-message") @@ -570,7 +570,7 @@ async def test_runnable_agent_with_function_calls() -> None: ], ) - def fake_parse(_: dict) -> AgentFinish | AgentAction: + def fake_parse(_: AIMessage) -> AgentFinish | AgentAction: """A parser.""" return cast("AgentFinish | AgentAction", next(parser_responses)) @@ -682,7 +682,7 @@ async def test_runnable_with_multi_action_per_step() -> None: ], ) - def fake_parse(_: dict) -> AgentFinish | AgentAction: + def fake_parse(_: AIMessage) -> AgentFinish | AgentAction: """A parser.""" return cast("AgentFinish | AgentAction", next(parser_responses))