chore(core): improve typing of Runnable __or__ (#34530)

`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 <github@mdrxy.com>
This commit is contained in:
Christophe BornetandMason Daugherty authored and GitHub committed 2026-06-10 16:16:03 -04:00
1 parent a063ec26dd
commit 720dfd3b09
8 files changed
+273 -114

No files matched your search

+37 -7
View File
@@ -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(
+146 -34
View File
@@ -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)}"
+2 -2
View File
@@ -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__(
@@ -233,7 +233,7 @@ def test_graph_sequence_map(snapshot: SnapshotAssertion) -> None:
return str_parser
return xml_parser
sequence: Runnable = (
sequence = (
prompt
| fake_llm
| {
@@ -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",
@@ -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}
@@ -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)
@@ -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))