mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
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:
1 parent
a063ec26dd
commit
720dfd3b09
8 files changed
+273
-114
No files matched your search
@@ -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(
|
||||
|
||||
@@ -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)}"
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
Reference in new issue
Block a user