chore(core): fix some any generics (#34545)

Co-authored-by: Mason Daugherty <github@mdrxy.com>
This commit is contained in:
Christophe BornetandMason Daugherty authored and GitHub committed 2026-06-10 15:32:14 -04:00
1 parent 3eee4002d9
commit a063ec26dd
100 files changed
+689 -531

No files matched your search

+2 -2
View File
@@ -331,7 +331,7 @@ def deprecated(
if not _obj_type:
_obj_type = "attribute"
wrapped = None
_name = _name or cast("type | Callable", obj.fget).__qualname__
_name = _name or cast("type", obj.fget).__qualname__
old_doc = obj.__doc__
class _DeprecatedProperty(property):
@@ -386,7 +386,7 @@ def deprecated(
return cast("T", prop)
else:
_name = _name or cast("type | Callable", obj).__qualname__
_name = _name or cast("type", obj).__qualname__
if not _obj_type:
# edge case: when a function is within another function
# within a test, this will call it a "method" not a "function"
+6 -4
View File
@@ -51,7 +51,7 @@ class AgentAction(Serializable):
tool: str
"""The name of the `Tool` to execute."""
tool_input: str | dict
tool_input: str | dict[Any, Any]
"""The input to pass in to the `Tool`."""
log: str
@@ -68,7 +68,9 @@ class AgentAction(Serializable):
type: Literal["AgentAction"] = "AgentAction"
# Override init to support instantiation by position for backward compat.
def __init__(self, tool: str, tool_input: str | dict, log: str, **kwargs: Any):
def __init__(
self, tool: str, tool_input: str | dict[Any, Any], log: str, **kwargs: Any
):
"""Create an `AgentAction`.
Args:
@@ -149,7 +151,7 @@ class AgentFinish(Serializable):
Agents return an `AgentFinish` when they have reached a stopping condition.
"""
return_values: dict
return_values: dict[Any, Any]
"""Dictionary of return values."""
log: str
@@ -164,7 +166,7 @@ class AgentFinish(Serializable):
"""
type: Literal["AgentFinish"] = "AgentFinish"
def __init__(self, return_values: dict, log: str, **kwargs: Any):
def __init__(self, return_values: dict[Any, Any], log: str, **kwargs: Any):
"""Override init to support instantiation by position for backward compat."""
super().__init__(return_values=return_values, log=log, **kwargs)
@@ -214,7 +214,7 @@ async def atrace_as_chain_group(
await run_manager.on_chain_end({})
Func = TypeVar("Func", bound=Callable)
Func = TypeVar("Func", bound=Callable[..., Any])
def shielded(func: Func) -> Func:
@@ -327,9 +327,7 @@ def handle_event(
# running coroutine, which we cannot interrupt to run this one.
# The solution is to run the synchronous function on the globally shared
# thread pool executor to avoid blocking the main event loop.
_executor().submit(
cast("Callable", copy_context().run), _run_coros, coros
).result()
_executor().submit(copy_context().run, _run_coros, coros).result()
else:
# If there's no running loop, we can run the coroutines directly.
_run_coros(coros)
@@ -381,10 +379,7 @@ async def _ahandle_event_for_handler(
else:
await asyncio.get_event_loop().run_in_executor(
None,
cast(
"Callable",
functools.partial(copy_context().run, event, *args, **kwargs),
),
functools.partial(copy_context().run, event, *args, **kwargs),
)
except NotImplementedError as e:
if event_name == "on_chat_model_start":
+2 -2
View File
@@ -1,7 +1,7 @@
"""**Chat Sessions** are a collection of messages and function calls."""
from collections.abc import Sequence
from typing import TypedDict
from typing import Any, TypedDict
from langchain_core.messages import BaseMessage
@@ -15,5 +15,5 @@ class ChatSession(TypedDict, total=False):
messages: Sequence[BaseMessage]
"""A sequence of the LangChain chat messages loaded from the source."""
functions: Sequence[dict]
functions: Sequence[dict[str, Any]]
"""A sequence of the function calling specs for the messages."""
@@ -48,7 +48,7 @@ class LangSmithLoader(BaseLoader):
inline_s3_urls: bool = True,
offset: int = 0,
limit: int | None = None,
metadata: dict | None = None,
metadata: dict[str, Any] | None = None,
filter: str | None = None, # noqa: A002
content_key: str = "",
format_content: Callable[..., str] | None = None,
+3 -3
View File
@@ -52,7 +52,7 @@ class BaseMedia(Serializable):
as a UUID, but this will not be enforced.
"""
metadata: dict = Field(default_factory=dict)
metadata: dict[Any, Any] = Field(default_factory=dict)
"""Arbitrary metadata associated with the content."""
@@ -218,7 +218,7 @@ class Blob(BaseMedia):
encoding: str = "utf-8",
mime_type: str | None = None,
guess_type: bool = True,
metadata: dict | None = None,
metadata: dict[Any, Any] | None = None,
) -> Blob:
"""Load the blob from a path like object.
@@ -255,7 +255,7 @@ class Blob(BaseMedia):
encoding: str = "utf-8",
mime_type: str | None = None,
path: str | None = None,
metadata: dict | None = None,
metadata: dict[Any, Any] | None = None,
) -> Blob:
"""Initialize the `Blob` from in-memory data.
+1 -1
View File
@@ -7,7 +7,7 @@ from langchain_core import __version__
@lru_cache(maxsize=1)
def get_runtime_environment() -> dict:
def get_runtime_environment() -> dict[str, str]:
"""Get information about the LangChain runtime environment.
Returns:
@@ -34,7 +34,7 @@ class BaseExampleSelector(ABC):
return await run_in_executor(None, self.add_example, example)
@abstractmethod
def select_examples(self, input_variables: dict[str, str]) -> list[dict]:
def select_examples(self, input_variables: dict[str, str]) -> list[dict[str, Any]]:
"""Select which examples to use based on the inputs.
Args:
@@ -45,7 +45,9 @@ class BaseExampleSelector(ABC):
A list of examples.
"""
async def aselect_examples(self, input_variables: dict[str, str]) -> list[dict]:
async def aselect_examples(
self, input_variables: dict[str, str]
) -> list[dict[str, Any]]:
"""Async select which examples to use based on the inputs.
Args:
@@ -2,6 +2,7 @@
import re
from collections.abc import Callable
from typing import Any
from pydantic import BaseModel, Field, model_validator
from typing_extensions import Self
@@ -48,7 +49,7 @@ class LengthBasedExampleSelector(BaseExampleSelector, BaseModel):
```
"""
examples: list[dict]
examples: list[dict[str, Any]]
"""A list of the examples that the prompt template expects."""
example_prompt: PromptTemplate
@@ -92,12 +93,12 @@ class LengthBasedExampleSelector(BaseExampleSelector, BaseModel):
self.example_text_lengths = [self.get_text_length(eg) for eg in string_examples]
return self
def select_examples(self, input_variables: dict[str, str]) -> list[dict]:
def select_examples(self, input_variables: dict[str, str]) -> list[dict[str, Any]]:
"""Select which examples to use based on the input lengths.
Args:
input_variables: A dictionary with keys as input variables
and values as their values.
and values as their values.
Returns:
A list of examples to include in the prompt.
@@ -115,12 +116,14 @@ class LengthBasedExampleSelector(BaseExampleSelector, BaseModel):
i += 1
return examples
async def aselect_examples(self, input_variables: dict[str, str]) -> list[dict]:
async def aselect_examples(
self, input_variables: dict[str, str]
) -> list[dict[str, Any]]:
"""Async select which examples to use based on the input lengths.
Args:
input_variables: A dictionary with keys as input variables
and values as their values.
and values as their values.
Returns:
A list of examples to include in the prompt.
@@ -15,7 +15,7 @@ if TYPE_CHECKING:
from langchain_core.embeddings import Embeddings
def sorted_values(values: dict[str, str]) -> list[Any]:
def sorted_values(values: dict[str, str]) -> list[str]:
"""Return a list of values in dict sorted by key.
Args:
@@ -33,13 +33,17 @@ class _VectorStoreExampleSelector(BaseExampleSelector, BaseModel, ABC):
vectorstore: VectorStore
"""VectorStore that contains information about examples."""
k: int = 4
"""Number of examples to select."""
example_keys: list[str] | None = None
"""Optional keys to filter examples to."""
input_keys: list[str] | None = None
"""Optional keys to filter input to. If provided, the search is based on
the input variables instead of all variables."""
vectorstore_kwargs: dict[str, Any] | None = None
"""Extra arguments passed to similarity_search function of the `VectorStore`."""
@@ -54,7 +58,7 @@ class _VectorStoreExampleSelector(BaseExampleSelector, BaseModel, ABC):
return " ".join(sorted_values({key: example[key] for key in input_keys}))
return " ".join(sorted_values(example))
def _documents_to_examples(self, documents: list[Document]) -> list[dict]:
def _documents_to_examples(self, documents: list[Document]) -> list[dict[str, Any]]:
# Get the examples from the metadata.
# This assumes that examples are stored in metadata.
examples = [dict(e.metadata) for e in documents]
@@ -97,7 +101,7 @@ class _VectorStoreExampleSelector(BaseExampleSelector, BaseModel, ABC):
class SemanticSimilarityExampleSelector(_VectorStoreExampleSelector):
"""Select examples based on semantic similarity."""
def select_examples(self, input_variables: dict[str, str]) -> list[dict]:
def select_examples(self, input_variables: dict[str, str]) -> list[dict[str, Any]]:
"""Select examples based on semantic similarity.
Args:
@@ -115,7 +119,9 @@ class SemanticSimilarityExampleSelector(_VectorStoreExampleSelector):
)
return self._documents_to_examples(example_docs)
async def aselect_examples(self, input_variables: dict[str, str]) -> list[dict]:
async def aselect_examples(
self, input_variables: dict[str, str]
) -> list[dict[str, Any]]:
"""Asynchronously select examples based on semantic similarity.
Args:
@@ -136,14 +142,14 @@ class SemanticSimilarityExampleSelector(_VectorStoreExampleSelector):
@classmethod
def from_examples(
cls,
examples: list[dict],
examples: list[dict[str, str]],
embeddings: Embeddings,
vectorstore_cls: type[VectorStore],
k: int = 4,
input_keys: list[str] | None = None,
*,
example_keys: list[str] | None = None,
vectorstore_kwargs: dict | None = None,
vectorstore_kwargs: dict[str, Any] | None = None,
**vectorstore_cls_kwargs: Any,
) -> SemanticSimilarityExampleSelector:
"""Create k-shot example selector using example list and embeddings.
@@ -180,14 +186,14 @@ class SemanticSimilarityExampleSelector(_VectorStoreExampleSelector):
@classmethod
async def afrom_examples(
cls,
examples: list[dict],
examples: list[dict[str, str]],
embeddings: Embeddings,
vectorstore_cls: type[VectorStore],
k: int = 4,
input_keys: list[str] | None = None,
*,
example_keys: list[str] | None = None,
vectorstore_kwargs: dict | None = None,
vectorstore_kwargs: dict[str, Any] | None = None,
**vectorstore_cls_kwargs: Any,
) -> SemanticSimilarityExampleSelector:
"""Async create k-shot example selector using example list and embeddings.
@@ -232,7 +238,7 @@ class MaxMarginalRelevanceExampleSelector(_VectorStoreExampleSelector):
fetch_k: int = 20
"""Number of examples to fetch to rerank."""
def select_examples(self, input_variables: dict[str, str]) -> list[dict]:
def select_examples(self, input_variables: dict[str, str]) -> list[dict[str, Any]]:
"""Select examples based on Max Marginal Relevance.
Args:
@@ -248,7 +254,9 @@ class MaxMarginalRelevanceExampleSelector(_VectorStoreExampleSelector):
)
return self._documents_to_examples(example_docs)
async def aselect_examples(self, input_variables: dict[str, str]) -> list[dict]:
async def aselect_examples(
self, input_variables: dict[str, str]
) -> list[dict[str, Any]]:
"""Asynchronously select examples based on Max Marginal Relevance.
Args:
@@ -267,14 +275,14 @@ class MaxMarginalRelevanceExampleSelector(_VectorStoreExampleSelector):
@classmethod
def from_examples(
cls,
examples: list[dict],
examples: list[dict[str, str]],
embeddings: Embeddings,
vectorstore_cls: type[VectorStore],
k: int = 4,
input_keys: list[str] | None = None,
fetch_k: int = 20,
example_keys: list[str] | None = None,
vectorstore_kwargs: dict | None = None,
vectorstore_kwargs: dict[str, Any] | None = None,
**vectorstore_cls_kwargs: Any,
) -> MaxMarginalRelevanceExampleSelector:
"""Create k-shot example selector using example list and embeddings.
@@ -313,7 +321,7 @@ class MaxMarginalRelevanceExampleSelector(_VectorStoreExampleSelector):
@classmethod
async def afrom_examples(
cls,
examples: list[dict],
examples: list[dict[str, str]],
embeddings: Embeddings,
vectorstore_cls: type[VectorStore],
*,
@@ -321,7 +329,7 @@ class MaxMarginalRelevanceExampleSelector(_VectorStoreExampleSelector):
input_keys: list[str] | None = None,
fetch_k: int = 20,
example_keys: list[str] | None = None,
vectorstore_kwargs: dict | None = None,
vectorstore_kwargs: dict[str, Any] | None = None,
**vectorstore_cls_kwargs: Any,
) -> MaxMarginalRelevanceExampleSelector:
"""Create k-shot example selector using example list and embeddings.
@@ -31,7 +31,7 @@ def _filter_invocation_params_for_tracing(params: dict[str, Any]) -> dict[str, A
def is_openai_data_block(
block: dict, filter_: Literal["image", "audio", "file"] | None = None
block: dict[str, Any], filter_: Literal["image", "audio", "file"] | None = None
) -> bool:
"""Check whether a block contains multimodal data in OpenAI Chat Completions format.
@@ -318,7 +318,7 @@ def _ensure_message_copy(message: T, formatted_message: T) -> T:
def _update_content_block(
formatted_message: "BaseMessage", idx: int, new_block: ContentBlock | dict
formatted_message: "BaseMessage", idx: int, new_block: ContentBlock | dict[str, Any]
) -> None:
"""Update a content block at the given index, handling type issues."""
# Type ignore needed because:
@@ -294,8 +294,8 @@ class BaseLanguageModel(
"""
def with_structured_output(
self, schema: dict | type, **kwargs: Any
) -> Runnable[LanguageModelInput, dict | BaseModel]:
self, schema: dict[str, Any] | type, **kwargs: Any
) -> Runnable[LanguageModelInput, dict[str, Any] | BaseModel]:
"""Not implemented on this class."""
# Implement this on child class if there is a way of steering the model to
# generate responses that match a given schema.
@@ -356,7 +356,7 @@ class BaseLanguageModel(
def get_num_tokens_from_messages(
self,
messages: list[BaseMessage],
tools: Sequence | None = None,
tools: Sequence[Any] | None = None,
) -> int:
"""Get the number of tokens in the messages.
@@ -67,6 +67,7 @@ from langchain_core.messages.block_translators.openai import (
)
from langchain_core.output_parsers.openai_tools import (
JsonOutputKeyToolsParser,
JsonOutputToolsParser,
PydanticToolsParser,
)
from langchain_core.outputs import (
@@ -99,7 +100,6 @@ if TYPE_CHECKING:
from langchain_protocol.protocol import MessagesData
from langchain_core.output_parsers.base import OutputParserLike
from langchain_core.runnables import Runnable, RunnableConfig
from langchain_core.runnables.schema import StreamEvent
from langchain_core.tools import BaseTool
@@ -108,7 +108,7 @@ if TYPE_CHECKING:
def _generate_response_from_error(error: BaseException) -> list[ChatGeneration]:
if hasattr(error, "response"):
response = error.response
metadata: dict = {}
metadata: dict[str, Any] = {}
if hasattr(response, "json"):
try:
metadata["body"] = response.json()
@@ -248,7 +248,9 @@ async def agenerate_from_stream(
return await run_in_executor(None, generate_from_stream, iter(chunks))
def _format_ls_structured_output(ls_structured_output_format: dict | None) -> dict:
def _format_ls_structured_output(
ls_structured_output_format: dict[str, Any] | None,
) -> dict[str, Any]:
if ls_structured_output_format:
try:
ls_structured_output_format_dict = {
@@ -804,7 +806,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
and isinstance(chunk.message, AIMessageChunk)
and not chunk.message.chunk_position
):
empty_content: str | list = (
empty_content: str | list[str | dict[str, Any]] = (
"" if isinstance(chunk.message.content, str) else []
)
msg_chunk = AIMessageChunk(
@@ -938,7 +940,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
and isinstance(chunk.message, AIMessageChunk)
and not chunk.message.chunk_position
):
empty_content: str | list = (
empty_content: str | list[str | dict[str, Any]] = (
"" if isinstance(chunk.message.content, str) else []
)
msg_chunk = AIMessageChunk(
@@ -1372,11 +1374,13 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
# --- Custom methods ---
def _combine_llm_outputs(
self, _llm_outputs: list[builtins.dict | None], /
) -> builtins.dict:
self, _llm_outputs: list[builtins.dict[str, Any] | None], /
) -> builtins.dict[str, Any]:
return {}
def _convert_cached_generations(self, cache_val: list) -> list[ChatGeneration]:
def _convert_cached_generations(
self, cache_val: list[Generation]
) -> list[ChatGeneration]:
"""Convert cached Generation objects to ChatGeneration objects.
Handle case where cache contains Generation objects instead of
@@ -1466,7 +1470,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
self,
stop: list[str] | None = None,
**kwargs: Any,
) -> builtins.dict:
) -> builtins.dict[str, Any]:
params = self._dict_for_compat()
params["stop"] = stop
return {**params, **kwargs}
@@ -1980,7 +1984,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
and isinstance(chunk.message, AIMessageChunk)
and not chunk.message.chunk_position
):
empty_content: str | list = (
empty_content: str | list[str | dict[str, Any]] = (
"" if isinstance(chunk.message.content, str) else []
)
chunk = ChatGenerationChunk(
@@ -2137,7 +2141,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
and isinstance(chunk.message, AIMessageChunk)
and not chunk.message.chunk_position
):
empty_content: str | list = (
empty_content: str | list[str | dict[str, Any]] = (
"" if isinstance(chunk.message.content, str) else []
)
chunk = ChatGenerationChunk(
@@ -2339,7 +2343,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
def bind_tools(
self,
tools: Sequence[builtins.dict[str, Any] | type | Callable | BaseTool],
tools: Sequence[builtins.dict[str, Any] | type | Callable[..., Any] | BaseTool],
*,
tool_choice: str | None = None,
**kwargs: Any,
@@ -2519,8 +2523,9 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
"schema": schema,
},
)
output_parser: JsonOutputToolsParser
if isinstance(schema, type) and is_basemodel_subclass(schema):
output_parser: OutputParserLike = PydanticToolsParser(
output_parser = PydanticToolsParser(
tools=[cast("TypeBaseModel", schema)], first_tool_only=True
)
else:
@@ -2679,7 +2684,7 @@ class SimpleChatModel(BaseChatModel):
def _gen_info_and_msg_metadata(
generation: ChatGeneration | ChatGenerationChunk,
) -> dict:
) -> dict[str, Any]:
return {
**(generation.generation_info or {}),
**generation.message.response_metadata,
@@ -6,7 +6,7 @@ These are traditionally older models (newer models generally are chat models).
from __future__ import annotations
import asyncio
import builtins # noqa: TC003
import builtins
import functools
import inspect
import json
@@ -60,11 +60,12 @@ from langchain_core.runnables import RunnableConfig, ensure_config, get_config_l
from langchain_core.runnables.config import run_in_executor
if TYPE_CHECKING:
import builtins
import uuid
logger = logging.getLogger(__name__)
_background_tasks: set[asyncio.Task] = set()
_background_tasks: set[asyncio.Task[None]] = set()
@functools.lru_cache
@@ -159,7 +160,7 @@ def get_prompts(
params: dict[str, Any],
prompts: list[str],
cache: BaseCache | bool | None = None, # noqa: FBT001
) -> tuple[dict[int, list], str, list[int], list[str]]:
) -> tuple[dict[int, list[Generation]], str, list[int], list[str]]:
"""Get prompts that are already cached.
Args:
@@ -195,7 +196,7 @@ async def aget_prompts(
params: dict[str, Any],
prompts: list[str],
cache: BaseCache | bool | None = None, # noqa: FBT001
) -> tuple[dict[int, list], str, list[int], list[str]]:
) -> tuple[dict[int, list[Generation]], str, list[int], list[str]]:
"""Get prompts that are already cached. Async version.
Args:
@@ -228,12 +229,12 @@ async def aget_prompts(
def update_cache(
cache: BaseCache | bool | None, # noqa: FBT001
existing_prompts: dict[int, list],
existing_prompts: dict[int, list[Generation]],
llm_string: str,
missing_prompt_idxs: list[int],
new_results: LLMResult,
prompts: list[str],
) -> dict | None:
) -> dict[str, Any] | None:
"""Update the cache and get the LLM output.
Args:
@@ -261,12 +262,12 @@ def update_cache(
async def aupdate_cache(
cache: BaseCache | bool | None, # noqa: FBT001
existing_prompts: dict[int, list],
existing_prompts: dict[int, list[Generation]],
llm_string: str,
missing_prompt_idxs: list[int],
new_results: LLMResult,
prompts: list[str],
) -> dict | None:
) -> dict[str, Any] | None:
"""Update the cache and get the LLM output. Async version.
Args:
@@ -954,7 +955,8 @@ class BaseLLM(BaseLanguageModel[str], ABC):
callbacks = cast("list[Callbacks]", callbacks)
tags_list = cast("list[list[str] | None]", tags or ([None] * len(prompts)))
metadata_list = cast(
"list[dict[str, Any] | None]", metadata or ([{}] * len(prompts))
"list[builtins.dict[str, Any] | None]",
metadata or ([{}] * len(prompts)),
)
run_name_list = run_name or cast(
"list[str | None]", ([None] * len(prompts))
@@ -989,7 +991,7 @@ class BaseLLM(BaseLanguageModel[str], ABC):
self.verbose,
cast("list[str]", tags),
self.tags,
cast("dict[str, Any]", metadata),
cast("builtins.dict[str, Any]", metadata),
self.metadata,
langsmith_inheritable_metadata=_filter_invocation_params_for_tracing(
params
@@ -1074,8 +1076,8 @@ class BaseLLM(BaseLanguageModel[str], ABC):
@staticmethod
def _get_run_ids_list(
run_id: uuid.UUID | list[uuid.UUID | None] | None, prompts: list
) -> list:
run_id: uuid.UUID | list[uuid.UUID | None] | None, prompts: list[str]
) -> list[uuid.UUID | None]:
if run_id is None:
return [None] * len(prompts)
if isinstance(run_id, list):
@@ -1226,7 +1228,8 @@ class BaseLLM(BaseLanguageModel[str], ABC):
callbacks = cast("list[Callbacks]", callbacks)
tags_list = cast("list[list[str] | None]", tags or ([None] * len(prompts)))
metadata_list = cast(
"list[dict[str, Any] | None]", metadata or ([{}] * len(prompts))
"list[builtins.dict[str, Any] | None]",
metadata or ([{}] * len(prompts)),
)
run_name_list = run_name or cast(
"list[str | None]", ([None] * len(prompts))
@@ -1261,7 +1264,7 @@ class BaseLLM(BaseLanguageModel[str], ABC):
self.verbose,
cast("list[str]", tags),
self.tags,
cast("dict[str, Any]", metadata),
cast("builtins.dict[str, Any]", metadata),
self.metadata,
langsmith_inheritable_metadata=_filter_invocation_params_for_tracing(
params
@@ -165,7 +165,7 @@ class Serializable(BaseModel, ABC):
return {}
@property
def lc_attributes(self) -> dict:
def lc_attributes(self) -> dict[str, Any]:
"""List of attribute names that should be included in the serialized kwargs.
These attributes must be accepted by the constructor.
+11 -11
View File
@@ -185,21 +185,21 @@ class AIMessage(BaseMessage):
@overload
def __init__(
self,
content: str | list[str | dict],
content: str | list[str | dict[Any, Any]],
**kwargs: Any,
) -> None: ...
@overload
def __init__(
self,
content: str | list[str | dict] | None = None,
content: str | list[str | dict[Any, Any]] | None = None,
content_blocks: list[types.ContentBlock] | None = None,
**kwargs: Any,
) -> None: ...
def __init__(
self,
content: str | list[str | dict] | None = None,
content: str | list[str | dict[Any, Any]] | None = None,
content_blocks: list[types.ContentBlock] | None = None,
**kwargs: Any,
) -> None:
@@ -221,14 +221,14 @@ class AIMessage(BaseMessage):
kwargs["tool_calls"] = content_tool_calls
super().__init__(
content=cast("str | list[str | dict]", content_blocks),
content=cast("list[str | dict[Any, Any]]", content_blocks),
**kwargs,
)
else:
super().__init__(content=content, **kwargs)
@property
def lc_attributes(self) -> dict:
def lc_attributes(self) -> dict[str, Any]:
"""Attributes to be serialized.
Includes all attributes, even if they are derived from other initialization
@@ -301,7 +301,7 @@ class AIMessage(BaseMessage):
# TODO: remove this logic if possible, reducing breaking nature of changes
@model_validator(mode="before")
@classmethod
def _backwards_compat_tool_calls(cls, values: dict) -> Any:
def _backwards_compat_tool_calls(cls, values: dict[str, Any]) -> Any:
check_additional_kwargs = not any(
values.get(k)
for k in ("tool_calls", "invalid_tool_calls", "tool_call_chunks")
@@ -431,7 +431,7 @@ class AIMessageChunk(AIMessage, BaseMessageChunk):
@property
@override
def lc_attributes(self) -> dict:
def lc_attributes(self) -> dict[str, Any]:
return {
"tool_calls": self.tool_calls,
"invalid_tool_calls": self.invalid_tool_calls,
@@ -769,8 +769,8 @@ def add_usage(left: UsageMetadata | None, right: UsageMetadata | None) -> UsageM
**cast(
"UsageMetadata",
_dict_int_op(
cast("dict", left),
cast("dict", right),
cast("dict[str, Any]", left),
cast("dict[str, Any]", right),
operator.add,
),
)
@@ -832,8 +832,8 @@ def subtract_usage(
**cast(
"UsageMetadata",
_dict_int_op(
cast("dict", left),
cast("dict", right),
cast("dict[str, Any]", left),
cast("dict[str, Any]", right),
(lambda le, ri: max(le - ri, 0)),
),
)
+13 -13
View File
@@ -100,10 +100,10 @@ class BaseMessage(Serializable):
[`SystemMessage`][langchain.messages.SystemMessage].
"""
content: str | list[str | dict]
content: str | list[str | dict[Any, Any]]
"""The contents of the message."""
additional_kwargs: dict = Field(default_factory=dict)
additional_kwargs: dict[Any, Any] = Field(default_factory=dict)
"""Reserved for additional payload data associated with the message.
For example, for a message from an AI, this could include tool calls as
@@ -111,7 +111,7 @@ class BaseMessage(Serializable):
"""
response_metadata: dict = Field(default_factory=dict)
response_metadata: dict[Any, Any] = Field(default_factory=dict)
"""Examples: response headers, logprobs, token counts, model name."""
type: str
@@ -146,21 +146,21 @@ class BaseMessage(Serializable):
@overload
def __init__(
self,
content: str | list[str | dict],
content: str | list[str | dict[Any, Any]],
**kwargs: Any,
) -> None: ...
@overload
def __init__(
self,
content: str | list[str | dict] | None = None,
content: str | list[str | dict[Any, Any]] | None = None,
content_blocks: list[types.ContentBlock] | None = None,
**kwargs: Any,
) -> None: ...
def __init__(
self,
content: str | list[str | dict] | None = None,
content: str | list[str | dict[Any, Any]] | None = None,
content_blocks: list[types.ContentBlock] | None = None,
**kwargs: Any,
) -> None:
@@ -364,9 +364,9 @@ class BaseMessage(Serializable):
def merge_content(
first_content: str | list[str | dict],
*contents: str | list[str | dict],
) -> str | list[str | dict]:
first_content: str | list[str | dict[Any, Any]],
*contents: str | list[str | dict[Any, Any]],
) -> str | list[str | dict[Any, Any]]:
"""Merge multiple message contents.
Args:
@@ -377,7 +377,7 @@ def merge_content(
The merged content.
"""
merged: str | list[str | dict]
merged: str | list[str | dict[Any, Any]]
merged = "" if first_content is None else first_content
for content in contents:
@@ -391,7 +391,7 @@ def merge_content(
merged = [merged, *content]
elif isinstance(content, list):
# If both are lists
merged = merge_lists(cast("list", merged), content) # type: ignore[assignment]
merged = merge_lists(merged, content) # type: ignore[assignment]
# If the first content is a list, and the second content is a string
# If the last element of the first content is a string
# Add the second content to the last element
@@ -471,7 +471,7 @@ class BaseMessageChunk(BaseMessage):
raise TypeError(msg)
def message_to_dict(message: BaseMessage) -> dict:
def message_to_dict(message: BaseMessage) -> dict[str, Any]:
"""Convert a Message to a dictionary.
Args:
@@ -485,7 +485,7 @@ def message_to_dict(message: BaseMessage) -> dict:
return {"type": message.type, "data": message.model_dump()}
def messages_to_dict(messages: Sequence[BaseMessage]) -> list[dict]:
def messages_to_dict(messages: Sequence[BaseMessage]) -> list[dict[str, Any]]:
"""Convert a sequence of Messages to a list of dictionaries.
Args:
@@ -197,8 +197,9 @@ def _convert_citation_to_v1(citation: dict[str, Any]) -> types.Annotation:
def _convert_to_v1_from_anthropic(message: AIMessage) -> list[types.ContentBlock]:
"""Convert Anthropic message content to v1 format."""
content: list[str | dict[str, Any]]
if isinstance(message.content, str):
content: list[str | dict] = [{"type": "text", "text": message.content}]
content = [{"type": "text", "text": message.content}]
else:
content = message.content
@@ -27,7 +27,7 @@ def _populate_extras(
return standard_block
def _parse_code_json(s: str) -> dict:
def _parse_code_json(s: str) -> dict[str, Any]:
"""Extract Python code from Groq built-in tool content.
Extracts the value of the 'code' field from a string of the form:
@@ -44,8 +44,8 @@ def _convert_v0_multimodal_input_to_v1(
def _convert_legacy_v0_content_block_to_v1(
block: dict,
) -> types.ContentBlock | dict:
block: dict[str, Any],
) -> types.ContentBlock | dict[str, Any]:
"""Convert a LangChain v0 content block to v1 format.
Preserves unknown keys as extras to avoid data loss.
@@ -53,7 +53,9 @@ def _convert_legacy_v0_content_block_to_v1(
Returns the original block unchanged if it's not in v0 format.
"""
def _extract_v0_extras(block_dict: dict, known_keys: set[str]) -> dict[str, Any]:
def _extract_v0_extras(
block_dict: dict[str, Any], known_keys: set[str]
) -> dict[str, Any]:
"""Extract unknown keys from v0 block to preserve as extras.
Args:
@@ -19,7 +19,7 @@ if TYPE_CHECKING:
from langchain_core.messages import AIMessage
def convert_to_openai_image_block(block: dict[str, Any]) -> dict:
def convert_to_openai_image_block(block: dict[str, Any]) -> dict[str, Any]:
"""Convert `ImageContentBlock` to format expected by OpenAI Chat Completions.
Args:
@@ -56,8 +56,9 @@ def convert_to_openai_image_block(block: dict[str, Any]) -> dict:
def convert_to_openai_data_block(
block: dict, api: Literal["chat/completions", "responses"] = "chat/completions"
) -> dict:
block: dict[str, Any],
api: Literal["chat/completions", "responses"] = "chat/completions",
) -> dict[str, Any]:
"""Format standard data content block to format expected by OpenAI.
"Standard data content block" can include old-style LangChain v0 blocks
@@ -265,7 +266,7 @@ def _convert_to_v1_from_chat_completions_chunk(
def _convert_from_v1_to_chat_completions(message: AIMessage) -> AIMessage:
"""Convert a v1 message to the Chat Completions format."""
if isinstance(message.content, list):
new_content: list = []
new_content: list[Any] = []
for block in message.content:
if isinstance(block, dict):
block_type = block.get("type")
@@ -330,7 +331,7 @@ def _convert_from_v03_ai_message(message: AIMessage) -> AIMessage:
]
# Build a bucket for every known block type
buckets: dict[str, list] = {key: [] for key in content_order}
buckets: dict[str, list[Any]] = {key: [] for key in content_order}
unknown_blocks = []
# Reasoning
@@ -422,8 +423,8 @@ def _convert_from_v03_ai_message(message: AIMessage) -> AIMessage:
def _convert_openai_format_to_data_block(
block: dict,
) -> types.ContentBlock | dict[Any, Any]:
block: dict[str, Any],
) -> types.ContentBlock | dict[str, Any]:
"""Convert OpenAI image/audio/file content block to respective v1 multimodal block.
We expect that the incoming block is verified to be in OpenAI Chat Completions
@@ -439,7 +440,9 @@ def _convert_openai_format_to_data_block(
"""
# Extract extra keys to put them in `extras`
def _extract_extras(block_dict: dict, known_keys: set[str]) -> dict[str, Any]:
def _extract_extras(
block_dict: dict[str, Any], known_keys: set[str]
) -> dict[str, Any]:
"""Extract unknown keys from block to preserve as extras."""
return {k: v for k, v in block_dict.items() if k not in known_keys}
+1 -1
View File
@@ -905,7 +905,7 @@ def _get_data_content_block_types() -> tuple[str, ...]:
return tuple(data_block_types)
def is_data_content_block(block: dict) -> bool:
def is_data_content_block(block: dict[str, Any]) -> bool:
"""Check if the provided content block is a data content block.
Returns True for both v0 (old-style) and v1 (new-style) multimodal data blocks.
+4 -4
View File
@@ -32,28 +32,28 @@ class HumanMessage(BaseMessage):
@overload
def __init__(
self,
content: str | list[str | dict],
content: str | list[str | dict[Any, Any]],
**kwargs: Any,
) -> None: ...
@overload
def __init__(
self,
content: str | list[str | dict] | None = None,
content: str | list[str | dict[Any, Any]] | None = None,
content_blocks: list[types.ContentBlock] | None = None,
**kwargs: Any,
) -> None: ...
def __init__(
self,
content: str | list[str | dict] | None = None,
content: str | list[str | dict[Any, Any]] | None = None,
content_blocks: list[types.ContentBlock] | None = None,
**kwargs: Any,
) -> None:
"""Specify `content` as positional arg or `content_blocks` for typing."""
if content_blocks is not None:
super().__init__(
content=cast("str | list[str | dict]", content_blocks),
content=cast("list[str | dict[Any, Any]]", content_blocks),
**kwargs,
)
else:
+4 -4
View File
@@ -32,28 +32,28 @@ class SystemMessage(BaseMessage):
@overload
def __init__(
self,
content: str | list[str | dict],
content: str | list[str | dict[Any, Any]],
**kwargs: Any,
) -> None: ...
@overload
def __init__(
self,
content: str | list[str | dict] | None = None,
content: str | list[str | dict[Any, Any]] | None = None,
content_blocks: list[types.ContentBlock] | None = None,
**kwargs: Any,
) -> None: ...
def __init__(
self,
content: str | list[str | dict] | None = None,
content: str | list[str | dict[Any, Any]] | None = None,
content_blocks: list[types.ContentBlock] | None = None,
**kwargs: Any,
) -> None:
"""Specify `content` as positional arg or `content_blocks` for typing."""
if content_blocks is not None:
super().__init__(
content=cast("str | list[str | dict]", content_blocks),
content=cast("list[str | dict[Any, Any]]", content_blocks),
**kwargs,
)
else:
+12 -10
View File
@@ -76,20 +76,20 @@ class ToolMessage(BaseMessage, ToolOutputMixin):
Should only be specified if it is different from the message content, e.g. if only
a subset of the full tool output is being passed as message content but the full
output is needed in other parts of the code.
"""
status: Literal["success", "error"] = "success"
"""Status of the tool invocation."""
additional_kwargs: dict = Field(default_factory=dict, repr=False)
additional_kwargs: dict[Any, Any] = Field(default_factory=dict, repr=False)
"""Currently inherited from `BaseMessage`, but not used."""
response_metadata: dict = Field(default_factory=dict, repr=False)
response_metadata: dict[Any, Any] = Field(default_factory=dict, repr=False)
"""Currently inherited from `BaseMessage`, but not used."""
@model_validator(mode="before")
@classmethod
def coerce_args(cls, values: dict) -> dict:
def coerce_args(cls, values: dict[str, Any]) -> dict[str, Any]:
"""Coerce the model arguments to the correct types.
Args:
@@ -135,21 +135,21 @@ class ToolMessage(BaseMessage, ToolOutputMixin):
@overload
def __init__(
self,
content: str | list[str | dict],
content: str | list[str | dict[Any, Any]],
**kwargs: Any,
) -> None: ...
@overload
def __init__(
self,
content: str | list[str | dict] | None = None,
content: str | list[str | dict[Any, Any]] | None = None,
content_blocks: list[types.ContentBlock] | None = None,
**kwargs: Any,
) -> None: ...
def __init__(
self,
content: str | list[str | dict] | None = None,
content: str | list[str | dict[Any, Any]] | None = None,
content_blocks: list[types.ContentBlock] | None = None,
**kwargs: Any,
) -> None:
@@ -164,7 +164,7 @@ class ToolMessage(BaseMessage, ToolOutputMixin):
"""
if content_blocks is not None:
super().__init__(
content=cast("str | list[str | dict]", content_blocks),
content=cast("list[str | dict[Any, Any]]", content_blocks),
**kwargs,
)
else:
@@ -347,7 +347,7 @@ def invalid_tool_call(
def default_tool_parser(
raw_tool_calls: list[dict],
raw_tool_calls: list[dict[str, Any]],
) -> tuple[list[ToolCall], list[InvalidToolCall]]:
"""Best-effort parsing of tools.
@@ -383,7 +383,9 @@ def default_tool_parser(
return tool_calls, invalid_tool_calls
def default_tool_chunk_parser(raw_tool_calls: list[dict]) -> list[ToolCallChunk]:
def default_tool_chunk_parser(
raw_tool_calls: list[dict[str, Any]],
) -> list[ToolCallChunk]:
"""Best-effort parsing of tool chunks.
Args:
+17 -16
View File
@@ -512,7 +512,7 @@ def get_buffer_string(
return message_separator.join(string_messages)
def _message_from_dict(message: dict) -> BaseMessage:
def _message_from_dict(message: dict[str, Any]) -> BaseMessage:
type_ = message["type"]
if type_ == "human":
return HumanMessage(**message["data"])
@@ -544,7 +544,7 @@ def _message_from_dict(message: dict) -> BaseMessage:
raise ValueError(msg)
def messages_from_dict(messages: Sequence[dict]) -> list[BaseMessage]:
def messages_from_dict(messages: Sequence[dict[str, Any]]) -> list[BaseMessage]:
"""Convert a sequence of messages from dicts to `Message` objects.
Args:
@@ -709,9 +709,9 @@ def _convert_to_message(message: MessageLikeRepresentation) -> BaseMessage:
- 2-tuple of (role string, template); e.g., (`'human'`, `'{user_input}'`)
- dict: a message dict with role and content keys
- dict: the `Serializable` constructor-envelope wire shape
`{"lc": 1, "type": "constructor", "id": [..., "<ClassName>"],
"kwargs": {...}}` — unpacked structurally and routed through the
standard dict-with-type dispatch.
`{"lc": 1, "type": "constructor", "id": [..., "<ClassName>"],
"kwargs": {...}}` — unpacked structurally and routed through the
standard dict-with-type dispatch.
- string: shorthand for (`'human'`, template); e.g., `'{user_input}'`
Args:
@@ -1132,7 +1132,7 @@ def trim_messages(
max_tokens: int,
token_counter: Callable[[list[BaseMessage]], int]
| Callable[[BaseMessage], int]
| BaseLanguageModel
| BaseLanguageModel[Any]
| Literal["approximate"],
strategy: Literal["first", "last"] = "last",
allow_partial: bool = False,
@@ -1484,10 +1484,11 @@ def trim_messages(
)
raise ValueError(msg)
text_splitter_fn: Callable[[str], list[str]]
if _HAS_LANGCHAIN_TEXT_SPLITTERS and isinstance(text_splitter, TextSplitter):
text_splitter_fn = text_splitter.split_text
elif text_splitter:
text_splitter_fn = cast("Callable", text_splitter)
text_splitter_fn = cast("Callable[[str], list[str]]", text_splitter)
else:
text_splitter_fn = _default_text_splitter
@@ -1528,7 +1529,7 @@ def convert_to_openai_messages(
text_format: Literal["string", "block"] = "string",
include_id: bool = False,
pass_through_unknown_blocks: bool = True,
) -> dict: ...
) -> dict[str, Any]: ...
@overload
@@ -1538,7 +1539,7 @@ def convert_to_openai_messages(
text_format: Literal["string", "block"] = "string",
include_id: bool = False,
pass_through_unknown_blocks: bool = True,
) -> list[dict]: ...
) -> list[dict[str, Any]]: ...
def convert_to_openai_messages(
@@ -1547,7 +1548,7 @@ def convert_to_openai_messages(
text_format: Literal["string", "block"] = "string",
include_id: bool = False,
pass_through_unknown_blocks: bool = True,
) -> dict | list[dict]:
) -> dict[str, Any] | list[dict[str, Any]]:
"""Convert LangChain messages into OpenAI message dicts.
Args:
@@ -1636,7 +1637,7 @@ def convert_to_openai_messages(
err = f"Unrecognized {text_format=}, expected one of 'string' or 'block'."
raise ValueError(err)
oai_messages: list[dict] = []
oai_messages: list[dict[str, Any]] = []
if is_single := isinstance(messages, (BaseMessage, dict, str)):
messages = [messages]
@@ -1644,9 +1645,9 @@ def convert_to_openai_messages(
messages = convert_to_messages(messages)
for i, message in enumerate(messages):
oai_msg: dict = {"role": _get_message_openai_role(message)}
tool_messages: list = []
content: str | list[dict]
oai_msg: dict[str, Any] = {"role": _get_message_openai_role(message)}
tool_messages: list[dict[str, Any]] = []
content: str | list[dict[str, Any]]
if message.name:
oai_msg["name"] = message.name
@@ -2216,7 +2217,7 @@ def _get_message_openai_role(message: BaseMessage) -> str:
raise ValueError(msg)
def _convert_to_openai_tool_calls(tool_calls: list[ToolCall]) -> list[dict]:
def _convert_to_openai_tool_calls(tool_calls: list[ToolCall]) -> list[dict[str, Any]]:
return [
{
"type": "function",
@@ -2247,7 +2248,7 @@ def count_tokens_approximately(
- For AI messages, the token count also includes stringified tool calls.
- For tool messages, the token count also includes the tool call ID.
- For multimodal messages with images, applies a fixed token penalty per image
instead of counting base64-encoded characters.
instead of counting base64-encoded characters.
- If tools are provided, the token count also includes stringified tool schemas.
Args:
@@ -2,7 +2,7 @@
from __future__ import annotations
import builtins # noqa: TC003
import builtins
import contextlib
from abc import ABC, abstractmethod
from typing import (
@@ -23,6 +23,8 @@ from langchain_core.runnables import Runnable, RunnableConfig, RunnableSerializa
from langchain_core.runnables.config import run_in_executor
if TYPE_CHECKING:
import builtins
from langchain_core.prompt_values import PromptValue
T = TypeVar("T")
@@ -70,7 +72,7 @@ class BaseLLMOutputParser(ABC, Generic[T]):
class BaseGenerationOutputParser(
BaseLLMOutputParser, RunnableSerializable[LanguageModelOutput, T]
BaseLLMOutputParser[T], RunnableSerializable[LanguageModelOutput, T]
):
"""Base class to parse the output of an LLM call."""
@@ -136,7 +138,7 @@ class BaseGenerationOutputParser(
class BaseOutputParser(
BaseLLMOutputParser, RunnableSerializable[LanguageModelOutput, T]
BaseLLMOutputParser[T], RunnableSerializable[LanguageModelOutput, T]
):
"""Base class to parse the output of an LLM call.
@@ -58,7 +58,7 @@ class ListOutputParser(BaseTransformOutputParser[list[str]]):
A list of strings.
"""
def parse_iter(self, text: str) -> Iterator[re.Match]:
def parse_iter(self, text: str) -> Iterator[re.Match[str]]:
"""Parse the output of an LLM call.
Args:
@@ -210,7 +210,7 @@ class NumberedListOutputParser(ListOutputParser):
return re.findall(self.pattern, text)
@override
def parse_iter(self, text: str) -> Iterator[re.Match]:
def parse_iter(self, text: str) -> Iterator[re.Match[str]]:
return re.finditer(self.pattern, text)
@property
@@ -241,7 +241,7 @@ class MarkdownListOutputParser(ListOutputParser):
return re.findall(self.pattern, text, re.MULTILINE)
@override
def parse_iter(self, text: str) -> Iterator[re.Match]:
def parse_iter(self, text: str) -> Iterator[re.Match[str]]:
return re.finditer(self.pattern, text, re.MULTILINE)
@property
@@ -100,7 +100,7 @@ def make_invalid_tool_call(
def parse_tool_calls(
raw_tool_calls: list[dict],
raw_tool_calls: list[dict[str, Any]],
*,
partial: bool = False,
strict: bool = False,
@@ -1,7 +1,7 @@
"""Output parsers using Pydantic."""
import json
from typing import Annotated, Generic, Literal, overload
from typing import Annotated, Any, Generic, Literal, overload
import pydantic
from pydantic import SkipValidation
@@ -22,7 +22,7 @@ class PydanticOutputParser(JsonOutputParser, Generic[TBaseModel]):
pydantic_object: Annotated[type[TBaseModel], SkipValidation()]
"""The Pydantic model to parse."""
def _parse_obj(self, obj: dict) -> TBaseModel:
def _parse_obj(self, obj: Any) -> TBaseModel:
try:
if issubclass(self.pydantic_object, pydantic.BaseModel):
return self.pydantic_object.model_validate(obj)
@@ -35,7 +35,7 @@ class PydanticOutputParser(JsonOutputParser, Generic[TBaseModel]):
raise self._parser_exception(e, obj) from e
def _parser_exception(
self, e: Exception, json_object: dict
self, e: Exception, json_object: Any
) -> OutputParserException:
json_string = json.dumps(json_object, ensure_ascii=False)
name = self.pydantic_object.__name__
@@ -148,7 +148,7 @@ class _StreamingParser:
self.pull_parser.close()
class XMLOutputParser(BaseTransformOutputParser):
class XMLOutputParser(BaseTransformOutputParser[dict[str, Any]]):
"""Parse an output using xml format.
Returns a dictionary of tags.
@@ -170,7 +170,7 @@ class XMLOutputParser(BaseTransformOutputParser):
3. A badly-formatted XML instance (unexpected 'tag' element):
`'<foo>\n <tag>\n </tag>\n</foo>'`
"""
encoding_matcher: re.Pattern = re.compile(
encoding_matcher: re.Pattern[str] = re.compile(
r"<([^>]*encoding[^>]*)>\n(.*)", re.MULTILINE | re.DOTALL
)
@@ -272,13 +272,13 @@ class XMLOutputParser(BaseTransformOutputParser):
# If root text contains any non-whitespace character it
# returns {root.tag: root.text}
return {root.tag: root.text}
result: dict = {root.tag: []}
root_tag: list[Any] = []
for child in root:
if len(child) == 0:
result[root.tag].append({child.tag: child.text})
root_tag.append({child.tag: child.text})
else:
result[root.tag].append(self._root_to_dict(child))
return result
root_tag.append(self._root_to_dict(child))
return {root.tag: root_tag}
@property
def _type(self) -> str:
@@ -1,5 +1,7 @@
"""Chat result schema."""
from typing import Any
from pydantic import BaseModel
from langchain_core.outputs.chat_generation import ChatGeneration
@@ -25,7 +27,7 @@ class ChatResult(BaseModel):
input prompt.
"""
llm_output: dict | None = None
llm_output: dict[str, Any] | None = None
"""For arbitrary model provider-specific output.
This dictionary is a free-form dictionary that can contain any information that the
@@ -3,7 +3,7 @@
from __future__ import annotations
from copy import deepcopy
from typing import Literal
from typing import Any, Literal
from pydantic import BaseModel
@@ -37,7 +37,7 @@ class LLMResult(BaseModel):
chat message.
"""
llm_output: dict | None = None
llm_output: dict[str, Any] | None = None
"""For arbitrary model provider-specific output.
This dictionary is a free-form dictionary that can contain any information that the
+2 -2
View File
@@ -8,7 +8,7 @@ from __future__ import annotations
from abc import ABC, abstractmethod
from collections.abc import Sequence
from typing import Literal, cast
from typing import Any, Literal, cast
from typing_extensions import TypedDict
@@ -146,7 +146,7 @@ class ImagePromptValue(PromptValue):
def to_messages(self) -> list[BaseMessage]:
"""Return prompt (image URL) as messages."""
return [HumanMessage(content=[cast("dict", self.image_url)])]
return [HumanMessage(content=[cast("dict[str, Any]", self.image_url)])]
class ChatPromptValueConcrete(ChatPromptValue):
+18 -10
View File
@@ -36,7 +36,7 @@ FormatOutputType = TypeVar("FormatOutputType")
class BasePromptTemplate(
RunnableSerializable[dict, PromptValue], ABC, Generic[FormatOutputType]
RunnableSerializable[dict[str, Any], PromptValue], ABC, Generic[FormatOutputType]
):
"""Base class for all prompt templates, returning a prompt."""
@@ -155,7 +155,7 @@ class BasePromptTemplate(
field_definitions={**required_input_variables, **optional_input_variables},
)
def _validate_input(self, inner_input: Any) -> builtins.dict:
def _validate_input(self, inner_input: Any) -> builtins.dict[str, Any]:
if not isinstance(inner_input, dict):
if len(self.input_variables) == 1:
var_name = self.input_variables[0]
@@ -192,22 +192,23 @@ class BasePromptTemplate(
return inner_input_
def _format_prompt_with_error_handling(
self,
inner_input: builtins.dict,
self, inner_input: builtins.dict[str, Any]
) -> PromptValue:
inner_input_ = self._validate_input(inner_input)
return self.format_prompt(**inner_input_)
async def _aformat_prompt_with_error_handling(
self,
inner_input: builtins.dict,
self, inner_input: builtins.dict[str, Any]
) -> PromptValue:
inner_input_ = self._validate_input(inner_input)
return await self.aformat_prompt(**inner_input_)
@override
def invoke(
self, input: builtins.dict, config: RunnableConfig | None = None, **kwargs: Any
self,
input: builtins.dict[str, Any],
config: RunnableConfig | None = None,
**kwargs: Any,
) -> PromptValue:
"""Invoke the prompt.
@@ -233,7 +234,10 @@ class BasePromptTemplate(
@override
async def ainvoke(
self, input: builtins.dict, config: RunnableConfig | None = None, **kwargs: Any
self,
input: builtins.dict[str, Any],
config: RunnableConfig | None = None,
**kwargs: Any,
) -> PromptValue:
"""Async invoke the prompt.
@@ -279,7 +283,9 @@ class BasePromptTemplate(
"""
return self.format_prompt(**kwargs)
def partial(self, **kwargs: str | Callable[[], str]) -> BasePromptTemplate:
def partial(
self, **kwargs: str | Callable[[], str]
) -> BasePromptTemplate[FormatOutputType]:
"""Return a partial of the prompt template.
Args:
@@ -420,7 +426,9 @@ class BasePromptTemplate(
raise ValueError(msg)
def _get_document_info(doc: Document, prompt: BasePromptTemplate[str]) -> dict:
def _get_document_info(
doc: Document, prompt: BasePromptTemplate[str]
) -> dict[str, Any]:
base_info = {"page_content": doc.page_content, **doc.metadata}
missing_metadata = set(prompt.input_variables).difference(base_info)
if len(missing_metadata) > 0:
+26 -22
View File
@@ -33,7 +33,7 @@ from langchain_core.messages import (
convert_to_messages,
)
from langchain_core.messages.base import get_msg_title_repr
from langchain_core.prompt_values import ChatPromptValue, ImageURL
from langchain_core.prompt_values import ChatPromptValue
from langchain_core.prompts.base import BasePromptTemplate
from langchain_core.prompts.dict import DictPromptTemplate
from langchain_core.prompts.image import ImagePromptTemplate
@@ -229,7 +229,7 @@ class BaseStringMessagePromptTemplate(BaseMessagePromptTemplate, ABC):
prompt: StringPromptTemplate
"""String prompt template."""
additional_kwargs: dict = Field(default_factory=dict)
additional_kwargs: dict[str, Any] = Field(default_factory=dict)
"""Additional keyword arguments to pass to the prompt template."""
@classmethod
@@ -387,11 +387,11 @@ class ChatMessagePromptTemplate(BaseStringMessagePromptTemplate):
class _TextTemplateParam(TypedDict, total=False):
text: str | dict
text: str | dict[str, Any]
class _ImageTemplateParam(TypedDict, total=False):
image_url: str | dict
image_url: str | dict[str, Any]
class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
@@ -402,7 +402,7 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
| list[StringPromptTemplate | ImagePromptTemplate | DictPromptTemplate]
)
"""Prompt template."""
additional_kwargs: dict = Field(default_factory=dict)
additional_kwargs: dict[str, Any] = Field(default_factory=dict)
"""Additional keyword arguments to pass to the prompt template."""
_msg_class: type[BaseMessage]
@@ -411,7 +411,7 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
def from_template(
cls: type[Self],
template: str
| list[str | _TextTemplateParam | _ImageTemplateParam | dict[str, Any]],
| Sequence[str | _TextTemplateParam | _ImageTemplateParam | dict[str, Any]],
template_format: PromptTemplateFormat = "f-string",
*,
partial_variables: dict[str, Any] | None = None,
@@ -434,14 +434,18 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
Raises:
ValueError: If the template is not a string or list of strings.
"""
prompt: (
StringPromptTemplate
| list[StringPromptTemplate | ImagePromptTemplate | DictPromptTemplate]
)
if isinstance(template, str):
prompt: StringPromptTemplate | list = PromptTemplate.from_template(
prompt = PromptTemplate.from_template(
template,
template_format=template_format,
partial_variables=partial_variables,
)
return cls(prompt=prompt, **kwargs)
if isinstance(template, list):
if isinstance(template, Sequence):
if (partial_variables is not None) and len(partial_variables) > 0:
msg = "Partial variables are not supported for list of templates."
raise ValueError(msg)
@@ -595,18 +599,18 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
return self._msg_class(
content=text, additional_kwargs=self.additional_kwargs
)
content: list = []
content: list[str | dict[str, Any]] = []
for prompt in self.prompt:
inputs = {var: kwargs[var] for var in prompt.input_variables}
if isinstance(prompt, StringPromptTemplate):
formatted_text: str = prompt.format(**inputs)
formatted_text = prompt.format(**inputs)
if formatted_text != "":
content.append({"type": "text", "text": formatted_text})
elif isinstance(prompt, ImagePromptTemplate):
formatted_image: ImageURL = prompt.format(**inputs)
formatted_image = prompt.format(**inputs)
content.append({"type": "image_url", "image_url": formatted_image})
elif isinstance(prompt, DictPromptTemplate):
formatted_dict: dict[str, Any] = prompt.format(**inputs)
formatted_dict = prompt.format(**inputs)
content.append(formatted_dict)
return self._msg_class(
content=content, additional_kwargs=self.additional_kwargs
@@ -626,18 +630,18 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
return self._msg_class(
content=text, additional_kwargs=self.additional_kwargs
)
content: list = []
content: list[str | dict[str, Any]] = []
for prompt in self.prompt:
inputs = {var: kwargs[var] for var in prompt.input_variables}
if isinstance(prompt, StringPromptTemplate):
formatted_text: str = await prompt.aformat(**inputs)
formatted_text = await prompt.aformat(**inputs)
if formatted_text != "":
content.append({"type": "text", "text": formatted_text})
elif isinstance(prompt, ImagePromptTemplate):
formatted_image: ImageURL = await prompt.aformat(**inputs)
formatted_image = await prompt.aformat(**inputs)
content.append({"type": "image_url", "image_url": formatted_image})
elif isinstance(prompt, DictPromptTemplate):
formatted_dict: dict[str, Any] = prompt.format(**inputs)
formatted_dict = prompt.format(**inputs)
content.append(formatted_dict)
return self._msg_class(
content=content, additional_kwargs=self.additional_kwargs
@@ -688,12 +692,12 @@ class SystemMessagePromptTemplate(_StringImageMessagePromptTemplate):
_msg_class: type[BaseMessage] = SystemMessage
class BaseChatPromptTemplate(BasePromptTemplate, ABC):
class BaseChatPromptTemplate(BasePromptTemplate[str], ABC):
"""Base class for chat prompt templates."""
@property
@override
def lc_attributes(self) -> dict:
def lc_attributes(self) -> dict[str, Any]:
return {"input_variables": self.input_variables}
def format(self, **kwargs: Any) -> str:
@@ -781,7 +785,7 @@ MessageLike = BaseMessagePromptTemplate | BaseMessage | BaseChatPromptTemplate
MessageLikeRepresentation = (
MessageLike
| tuple[str | type, str | Sequence[dict] | Sequence[object]]
| tuple[str | type, str | Sequence[dict[str, Any]] | Sequence[object]]
| str
| dict[str, Any]
)
@@ -1046,7 +1050,7 @@ class ChatPromptTemplate(BaseChatPromptTemplate):
@model_validator(mode="before")
@classmethod
def validate_input_variables(cls, values: dict) -> Any:
def validate_input_variables(cls, values: dict[str, Any]) -> Any:
"""Validate input variables.
If `input_variables` is not set, it will be set to the union of all input
@@ -1062,7 +1066,7 @@ class ChatPromptTemplate(BaseChatPromptTemplate):
ValueError: If input variables do not match.
"""
messages = values["messages"]
input_vars: set = set()
input_vars: set[str] = set()
optional_variables = set()
input_types: dict[str, Any] = values.get("input_types", {})
for message in messages:
@@ -1336,7 +1340,7 @@ class ChatPromptTemplate(BaseChatPromptTemplate):
def _create_template_from_message_type(
message_type: str,
template: str | list,
template: str | list[str | dict[str, Any]],
template_format: PromptTemplateFormat = "f-string",
) -> BaseMessagePromptTemplate:
"""Create a message prompt template from a message type and template string.
+4 -4
View File
@@ -16,7 +16,7 @@ from langchain_core.runnables import RunnableConfig, RunnableSerializable
from langchain_core.runnables.config import ensure_config
class DictPromptTemplate(RunnableSerializable[dict, dict]):
class DictPromptTemplate(RunnableSerializable[dict[str, Any], dict[str, Any]]):
"""Template represented by a dictionary.
Recognizes variables in f-string or mustache formatted string dict values.
@@ -74,8 +74,8 @@ class DictPromptTemplate(RunnableSerializable[dict, dict]):
@override
def invoke(
self, input: dict, config: RunnableConfig | None = None, **kwargs: Any
) -> dict:
self, input: dict[str, Any], config: RunnableConfig | None = None, **kwargs: Any
) -> dict[str, Any]:
return self._call_with_config(
lambda x: self.format(**x),
input,
@@ -123,7 +123,7 @@ class DictPromptTemplate(RunnableSerializable[dict, dict]):
def _get_input_variables(
template: dict, template_format: Literal["f-string", "mustache"]
template: dict[str, Any], template_format: Literal["f-string", "mustache"]
) -> list[str]:
input_variables = []
for v in template.values():
+4 -4
View File
@@ -34,7 +34,7 @@ if TYPE_CHECKING:
class _FewShotPromptTemplateMixin(BaseModel):
"""Prompt template that contains few shot examples."""
examples: list[dict] | None = None
examples: list[dict[str, Any]] | None = None
"""Examples to format into the prompt.
Either this or `example_selector` should be provided.
@@ -53,7 +53,7 @@ class _FewShotPromptTemplateMixin(BaseModel):
@model_validator(mode="before")
@classmethod
def check_examples_and_selector(cls, values: dict) -> Any:
def check_examples_and_selector(cls, values: dict[str, Any]) -> Any:
"""Check that one and only one of `examples`/`example_selector` are provided.
Args:
@@ -79,7 +79,7 @@ class _FewShotPromptTemplateMixin(BaseModel):
return values
def _get_examples(self, **kwargs: Any) -> list[dict]:
def _get_examples(self, **kwargs: Any) -> list[dict[str, Any]]:
"""Get the examples to use for formatting the prompt.
Args:
@@ -98,7 +98,7 @@ class _FewShotPromptTemplateMixin(BaseModel):
msg = "One of 'examples' and 'example_selector' should be provided"
raise ValueError(msg)
async def _aget_examples(self, **kwargs: Any) -> list[dict]:
async def _aget_examples(self, **kwargs: Any) -> list[dict[str, Any]]:
"""Async get the examples to use for formatting the prompt.
Args:
@@ -19,7 +19,7 @@ from langchain_core.prompts.string import (
class FewShotPromptWithTemplates(StringPromptTemplate):
"""Prompt template that contains few shot examples."""
examples: list[dict] | None = None
examples: list[dict[str, Any]] | None = None
"""Examples to format into the prompt.
Either this or `example_selector` should be provided.
@@ -63,7 +63,7 @@ class FewShotPromptWithTemplates(StringPromptTemplate):
@model_validator(mode="before")
@classmethod
def check_examples_and_selector(cls, values: dict) -> Any:
def check_examples_and_selector(cls, values: dict[str, Any]) -> Any:
"""Check that one and only one of examples/example_selector are provided."""
examples = values.get("examples")
example_selector = values.get("example_selector")
@@ -106,14 +106,14 @@ class FewShotPromptWithTemplates(StringPromptTemplate):
extra="forbid",
)
def _get_examples(self, **kwargs: Any) -> list[dict]:
def _get_examples(self, **kwargs: Any) -> list[dict[str, Any]]:
if self.examples is not None:
return self.examples
if self.example_selector is not None:
return self.example_selector.select_examples(kwargs)
raise ValueError
async def _aget_examples(self, **kwargs: Any) -> list[dict]:
async def _aget_examples(self, **kwargs: Any) -> list[dict[str, Any]]:
if self.examples is not None:
return self.examples
if self.example_selector is not None:
+1 -1
View File
@@ -29,7 +29,7 @@ class ImagePromptTemplate(BasePromptTemplate[ImageURL]):
```
"""
template: dict = Field(default_factory=dict)
template: dict[str, Any] = Field(default_factory=dict)
"""Template for the prompt."""
template_format: PromptTemplateFormat = "f-string"
+15 -12
View File
@@ -4,6 +4,7 @@ import json
import logging
from collections.abc import Callable
from pathlib import Path
from typing import Any
import yaml
@@ -52,8 +53,8 @@ def _validate_path(path: Path) -> None:
"prompts and `load`/`loads` to deserialize them.",
)
def load_prompt_from_config(
config: dict, *, allow_dangerous_paths: bool = False
) -> BasePromptTemplate:
config: dict[str, Any], *, allow_dangerous_paths: bool = False
) -> BasePromptTemplate[str]:
"""Load prompt from config dict.
Args:
@@ -83,8 +84,8 @@ def load_prompt_from_config(
def _load_template(
var_name: str, config: dict, *, allow_dangerous_paths: bool = False
) -> dict:
var_name: str, config: dict[str, Any], *, allow_dangerous_paths: bool = False
) -> dict[str, Any]:
"""Load template from the path if applicable."""
# Check if template_path exists in config.
if f"{var_name}_path" in config:
@@ -109,7 +110,9 @@ def _load_template(
return config
def _load_examples(config: dict, *, allow_dangerous_paths: bool = False) -> dict:
def _load_examples(
config: dict[str, Any], *, allow_dangerous_paths: bool = False
) -> dict[str, Any]:
"""Load examples if necessary."""
if isinstance(config["examples"], list):
pass
@@ -132,7 +135,7 @@ def _load_examples(config: dict, *, allow_dangerous_paths: bool = False) -> dict
return config
def _load_output_parser(config: dict) -> dict:
def _load_output_parser(config: dict[str, Any]) -> dict[str, Any]:
"""Load output parser."""
if config_ := config.get("output_parser"):
if output_parser_type := config_.get("_type") != "default":
@@ -143,7 +146,7 @@ def _load_output_parser(config: dict) -> dict:
def _load_few_shot_prompt(
config: dict, *, allow_dangerous_paths: bool = False
config: dict[str, Any], *, allow_dangerous_paths: bool = False
) -> FewShotPromptTemplate:
"""Load the "few shot" prompt from the config."""
# Load the suffix and prefix templates.
@@ -178,7 +181,7 @@ def _load_few_shot_prompt(
def _load_prompt(
config: dict, *, allow_dangerous_paths: bool = False
config: dict[str, Any], *, allow_dangerous_paths: bool = False
) -> PromptTemplate:
"""Load the prompt template from config."""
# Load the template from disk if necessary.
@@ -212,7 +215,7 @@ def load_prompt(
encoding: str | None = None,
*,
allow_dangerous_paths: bool = False,
) -> BasePromptTemplate:
) -> BasePromptTemplate[str]:
"""Unified method for loading a prompt from LangChainHub or local filesystem.
Args:
@@ -247,7 +250,7 @@ def _load_prompt_from_file(
encoding: str | None = None,
*,
allow_dangerous_paths: bool = False,
) -> BasePromptTemplate:
) -> BasePromptTemplate[str]:
"""Load prompt from file."""
# Convert file to a Path object.
file_path = Path(file)
@@ -266,7 +269,7 @@ def _load_prompt_from_file(
def _load_chat_prompt(
config: dict,
config: dict[str, Any],
*,
allow_dangerous_paths: bool = False, # noqa: ARG001
) -> ChatPromptTemplate:
@@ -282,7 +285,7 @@ def _load_chat_prompt(
return ChatPromptTemplate.from_template(template=template, **config)
type_to_loader_dict: dict[str, Callable[..., BasePromptTemplate]] = {
type_to_loader_dict: dict[str, Callable[..., BasePromptTemplate[str]]] = {
"prompt": _load_prompt,
"few_shot": _load_few_shot_prompt,
"chat": _load_chat_prompt,
+1 -1
View File
@@ -88,7 +88,7 @@ class PromptTemplate(StringPromptTemplate):
@model_validator(mode="before")
@classmethod
def pre_init_validation(cls, values: dict) -> Any:
def pre_init_validation(cls, values: dict[str, Any]) -> Any:
"""Check that template and input variables are consistent."""
if values.get("template") is None:
# Will let pydantic fail with a ValidationError if template
+3 -3
View File
@@ -213,7 +213,7 @@ DEFAULT_FORMATTER_MAPPING: dict[str, Callable[..., str]] = {
"jinja2": jinja2_formatter,
}
DEFAULT_VALIDATOR_MAPPING: dict[str, Callable] = {
DEFAULT_VALIDATOR_MAPPING: dict[str, Callable[[str, list[str]], None]] = {
"f-string": formatter.validate_input_variables,
"jinja2": validate_jinja2,
}
@@ -325,7 +325,7 @@ def get_template_variables(template: str, template_format: str) -> list[str]:
return sorted(input_variables)
class StringPromptTemplate(BasePromptTemplate, ABC):
class StringPromptTemplate(BasePromptTemplate[str], ABC):
"""String prompt that exposes the format method, returning a prompt."""
@classmethod
@@ -390,7 +390,7 @@ class StringPromptTemplate(BasePromptTemplate, ABC):
print(self.pretty_repr(html=is_interactive_env())) # noqa: T201
def is_subsequence(child: Sequence, parent: Sequence) -> bool:
def is_subsequence(child: Sequence[Any], parent: Sequence[Any]) -> bool:
"""Return `True` if child is subsequence of parent."""
if len(child) == 0 or len(parent) == 0:
return False
@@ -28,7 +28,7 @@ from langchain_core.utils import get_pydantic_field_names
class StructuredPrompt(ChatPromptTemplate):
"""Structured prompt template for a language model."""
schema_: dict | type
schema_: dict[str, Any] | type
"""Schema for the structured prompt."""
structured_output_kwargs: dict[str, Any] = Field(default_factory=dict)
@@ -36,7 +36,7 @@ class StructuredPrompt(ChatPromptTemplate):
def __init__(
self,
messages: Sequence[MessageLikeRepresentation],
schema_: dict | type[BaseModel] | None = None,
schema_: dict[str, Any] | type[BaseModel] | None = None,
*,
structured_output_kwargs: dict[str, Any] | None = None,
template_format: PromptTemplateFormat = "f-string",
@@ -87,7 +87,7 @@ class StructuredPrompt(ChatPromptTemplate):
def from_messages_and_schema(
cls,
messages: Sequence[MessageLikeRepresentation],
schema: dict | type,
schema: dict[str, Any] | type,
**kwargs: Any,
) -> ChatPromptTemplate:
"""Create a chat prompt template from a variety of message formats.
@@ -143,7 +143,7 @@ class StructuredPrompt(ChatPromptTemplate):
| Callable[[AsyncIterator[Any]], AsyncIterator[Other]]
| Callable[[Any], Other]
| Mapping[str, Runnable[Any, Other] | Callable[[Any], Other] | Any],
) -> RunnableSerializable[dict, Other]:
) -> RunnableSerializable[dict[str, Any], Other]:
return self.pipe(other)
def pipe(
@@ -154,7 +154,7 @@ class StructuredPrompt(ChatPromptTemplate):
| Callable[[Any], Other]
| Mapping[str, Runnable[Any, Other] | Callable[[Any], Other] | Any],
name: str | None = None,
) -> RunnableSerializable[dict, Other]:
) -> RunnableSerializable[dict[str, Any], Other]:
"""Pipe the structured prompt to a language model.
Args:
+17 -15
View File
@@ -605,7 +605,7 @@ class Runnable(ABC, Generic[Input, Output]):
def get_prompts(
self, config: RunnableConfig | None = None
) -> list[BasePromptTemplate]:
) -> list[BasePromptTemplate[Any]]:
"""Return a list of prompts used by this `Runnable`."""
# Import locally to prevent circular import
from langchain_core.prompts.base import BasePromptTemplate # noqa: PLC0415
@@ -2488,7 +2488,7 @@ class Runnable(ABC, Generic[Input, Output]):
iterator = context.run(transformer, input_for_transform, **kwargs) # type: ignore[arg-type]
if stream_handler := next(
(
cast("_StreamingCallbackHandler", h)
h
for h in run_manager.handlers
# instance check OK here, it's a mixin
if isinstance(h, _StreamingCallbackHandler)
@@ -2590,7 +2590,7 @@ class Runnable(ABC, Generic[Input, Output]):
if stream_handler := next(
(
cast("_StreamingCallbackHandler", h)
h
for h in run_manager.handlers
# instance check OK here, it's a mixin
if isinstance(h, _StreamingCallbackHandler)
@@ -3088,7 +3088,7 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
def __init__(
self,
*steps: RunnableLike,
*steps: RunnableLike[Any, Any],
name: str | None = None,
first: Runnable[Any, Any] | None = None,
middle: list[Runnable[Any, Any]] | None = None,
@@ -3106,7 +3106,7 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
Raises:
ValueError: If the sequence has less than 2 steps.
"""
steps_flat: list[Runnable] = []
steps_flat: list[Runnable[Any, Any]] = []
if not steps and first is not None and last is not None:
steps_flat = [first] + (middle or []) + [last]
for step in steps:
@@ -4216,7 +4216,7 @@ class RunnableParallel(RunnableSerializable[Input, dict[str, Any]]):
]
# Wrap in a coroutine to satisfy linter
async def get_next_chunk(generator: AsyncIterator) -> Output | None:
async def get_next_chunk(generator: AsyncIterator[Any]) -> Output | None:
return await anext(generator)
# Start the first iteration of each generator
@@ -4383,9 +4383,10 @@ class RunnableGenerator(Runnable[Input, Output]):
TypeError: If the transform is not a generator function.
"""
func_for_name: Callable[..., Any]
if atransform is not None:
self._atransform = atransform
func_for_name: Callable = atransform
func_for_name = atransform
if is_async_generator(transform):
self._atransform = transform
@@ -4795,9 +4796,10 @@ class RunnableLambda(Runnable[Input, Output]):
TypeError: If both `func` and `afunc` are provided.
"""
func_for_name: Callable[..., Any]
if afunc is not None:
self.afunc = afunc
func_for_name: Callable = afunc
func_for_name = afunc
if is_async_callable(func) or is_async_generator(func):
if afunc is not None:
@@ -4939,7 +4941,7 @@ class RunnableLambda(Runnable[Input, Output]):
)
@functools.cached_property
def deps(self) -> list[Runnable]:
def deps(self) -> list[Runnable[Any, Any]]:
"""The dependencies of this `Runnable`.
Returns:
@@ -4954,7 +4956,7 @@ class RunnableLambda(Runnable[Input, Output]):
else:
objects = []
deps: list[Runnable] = []
deps: list[Runnable[Any, Any]] = []
for obj in objects:
if isinstance(obj, Runnable):
deps.append(obj)
@@ -5130,7 +5132,7 @@ class RunnableLambda(Runnable[Input, Output]):
cast(
"AsyncGenerator[Any, Any]",
acall_func_with_variable_args(
cast("Callable", afunc),
cast("Callable[..., Any]", afunc),
value,
config,
run_manager,
@@ -5151,7 +5153,7 @@ class RunnableLambda(Runnable[Input, Output]):
output = chunk
else:
output = await acall_func_with_variable_args(
cast("Callable", afunc), value, config, run_manager, **kwargs
cast("Callable[..., Any]", afunc), value, config, run_manager, **kwargs
)
# If the output is a Runnable, invoke it
if isinstance(output, Runnable):
@@ -5373,7 +5375,7 @@ class RunnableLambda(Runnable[Input, Output]):
async for chunk in cast(
"AsyncIterator[Output]",
acall_func_with_variable_args(
cast("Callable", afunc),
cast("Callable[..., Any]", afunc),
final,
config,
run_manager,
@@ -5390,7 +5392,7 @@ class RunnableLambda(Runnable[Input, Output]):
output = chunk
else:
output = await acall_func_with_variable_args(
cast("Callable", afunc),
cast("Callable[..., Any]", afunc),
final,
config,
run_manager,
@@ -6486,7 +6488,7 @@ RunnableLike = (
)
def coerce_to_runnable(thing: RunnableLike) -> Runnable[Input, Output]:
def coerce_to_runnable(thing: RunnableLike[Input, Output]) -> Runnable[Input, Output]:
"""Coerce a `Runnable`-like object into a `Runnable`.
Args:
+3 -5
View File
@@ -77,9 +77,9 @@ class RunnableBranch(RunnableSerializable[Input, Output]):
Runnable[Input, bool]
| Callable[[Input], bool]
| Callable[[Input], Awaitable[bool]],
RunnableLike,
RunnableLike[Input, Output],
]
| RunnableLike,
| RunnableLike[Input, Output],
) -> None:
"""A `Runnable` that runs one of two branches based on a condition.
@@ -106,9 +106,7 @@ class RunnableBranch(RunnableSerializable[Input, Output]):
msg = "RunnableBranch default must be Runnable, callable or mapping."
raise TypeError(msg)
default_ = cast(
"Runnable[Input, Output]", coerce_to_runnable(cast("RunnableLike", default))
)
default_ = coerce_to_runnable(cast("RunnableLike[Input, Output]", default))
branches_ = []
+3 -3
View File
@@ -97,7 +97,7 @@ class Node(NamedTuple):
"""The unique identifier of the node."""
name: str
"""The name of the node."""
data: type[BaseModel] | RunnableType | None
data: type[BaseModel] | RunnableType[Any, Any] | None
"""The data of the node."""
metadata: dict[str, Any] | None
"""Optional metadata for the node. """
@@ -177,7 +177,7 @@ class MermaidDrawMethod(Enum):
def node_data_str(
id: str,
data: type[BaseModel] | RunnableType | None,
data: type[BaseModel] | RunnableType[Any, Any] | None,
) -> str:
"""Convert the data of a node to a string.
@@ -311,7 +311,7 @@ class Graph:
def add_node(
self,
data: type[BaseModel] | RunnableType | None,
data: type[BaseModel] | RunnableType[Any, Any] | None,
id: str | None = None,
*,
metadata: dict[str, Any] | None = None,
@@ -36,7 +36,7 @@ MessagesOrDictWithMessages = Sequence["BaseMessage"] | dict[str, Any]
GetSessionHistoryCallable = Callable[..., BaseChatMessageHistory]
class RunnableWithMessageHistory(RunnableBindingBase): # type: ignore[no-redef]
class RunnableWithMessageHistory(RunnableBindingBase[Any, Any]): # type: ignore[no-redef]
"""`Runnable` that manages chat message history for another `Runnable`.
A chat message history is a sequence of messages that represent a conversation.
@@ -391,7 +391,7 @@ class RunnableWithMessageHistory(RunnableBindingBase): # type: ignore[no-redef]
@override
def get_input_schema(self, config: RunnableConfig | None = None) -> type[BaseModel]:
fields: dict = {}
fields: dict[str, Any] = {}
if self.input_messages_key and self.history_messages_key:
fields[self.input_messages_key] = (
str | BaseMessage | Sequence[BaseMessage],
@@ -450,7 +450,7 @@ class RunnableWithMessageHistory(RunnableBindingBase): # type: ignore[no-redef]
)
def _get_input_messages(
self, input_val: str | BaseMessage | Sequence[BaseMessage] | dict
self, input_val: str | BaseMessage | Sequence[BaseMessage] | dict[str, Any]
) -> list[BaseMessage]:
# If dictionary, try to pluck the single key representing messages
if isinstance(input_val, dict):
@@ -488,7 +488,7 @@ class RunnableWithMessageHistory(RunnableBindingBase): # type: ignore[no-redef]
raise ValueError(msg)
def _get_output_messages(
self, output_val: str | BaseMessage | Sequence[BaseMessage] | dict
self, output_val: str | BaseMessage | Sequence[BaseMessage] | dict[str, Any]
) -> list[BaseMessage]:
# If dictionary, try to pluck the single key representing messages
if isinstance(output_val, dict):
@@ -9,7 +9,6 @@ from collections.abc import Awaitable, Callable
from typing import (
TYPE_CHECKING,
Any,
cast,
)
from pydantic import BaseModel, RootModel
@@ -346,7 +345,7 @@ class RunnablePassthrough(RunnableSerializable[Other, Other]):
yield chunk
_graph_passthrough: RunnablePassthrough = RunnablePassthrough()
_graph_passthrough = RunnablePassthrough[Any]()
class RunnableAssign(RunnableSerializable[dict[str, Any], dict[str, Any]]):
@@ -577,9 +576,11 @@ class RunnableAssign(RunnableSerializable[dict[str, Any], dict[str, Any]]):
if filtered:
yield filtered
# yield map output
yield cast("dict[str, Any]", first_map_chunk_future.result())
for chunk in map_output:
yield chunk
first_chunk = first_map_chunk_future.result()
if first_chunk is not None:
yield first_chunk
for chunk in map_output:
yield chunk
@override
def transform(
@@ -613,7 +614,7 @@ class RunnableAssign(RunnableSerializable[dict[str, Any], dict[str, Any]]):
**kwargs,
)
# start map output stream
first_map_chunk_task: asyncio.Task = asyncio.create_task(
first_map_chunk_task = asyncio.create_task(
anext(map_output, None),
)
# consume passthrough stream
@@ -629,9 +630,11 @@ class RunnableAssign(RunnableSerializable[dict[str, Any], dict[str, Any]]):
if filtered:
yield filtered
# yield map output
yield await first_map_chunk_task
async for chunk in map_output:
yield chunk
first_chunk = await first_map_chunk_task
if first_chunk is not None:
yield first_chunk
async for chunk in map_output:
yield chunk
@override
async def atransform(
+11 -7
View File
@@ -46,7 +46,9 @@ Input = TypeVar("Input", contravariant=True) # noqa: PLC0105
Output = TypeVar("Output", covariant=True) # noqa: PLC0105
async def gated_coro(semaphore: asyncio.Semaphore, coro: Coroutine) -> Any:
async def gated_coro(
semaphore: asyncio.Semaphore, coro: Coroutine[Any, Any, Any]
) -> Any:
"""Run a coroutine with a semaphore.
Args:
@@ -60,7 +62,9 @@ async def gated_coro(semaphore: asyncio.Semaphore, coro: Coroutine) -> Any:
return await coro
async def gather_with_concurrency(n: int | None, *coros: Coroutine) -> list:
async def gather_with_concurrency(
n: int | None, *coros: Coroutine[Any, Any, Any]
) -> list[Any]:
"""Gather coroutines with a limit on the number of concurrent coroutines.
Args:
@@ -362,7 +366,7 @@ class GetLambdaSource(ast.NodeVisitor):
self.source = ast.unparse(node)
def get_function_first_arg_dict_keys(func: Callable) -> list[str] | None:
def get_function_first_arg_dict_keys(func: Callable[..., Any]) -> list[str] | None:
"""Get the keys of the first argument of a function if it is a dict.
Args:
@@ -381,7 +385,7 @@ def get_function_first_arg_dict_keys(func: Callable) -> list[str] | None:
return None
def get_lambda_source(func: Callable) -> str | None:
def get_lambda_source(func: Callable[..., Any]) -> str | None:
"""Get the source code of a lambda function.
Args:
@@ -405,7 +409,7 @@ def get_lambda_source(func: Callable) -> str | None:
@lru_cache(maxsize=256)
def get_function_nonlocals(func: Callable) -> list[Any]:
def get_function_nonlocals(func: Callable[..., Any]) -> list[Any]:
"""Get the nonlocal variables accessed by a function.
Args:
@@ -747,7 +751,7 @@ class _RootEventFilter:
def is_async_generator(
func: Any,
) -> TypeGuard[Callable[..., AsyncIterator]]:
) -> TypeGuard[Callable[..., AsyncIterator[Any]]]:
"""Check if a function is an async generator.
Args:
@@ -764,7 +768,7 @@ def is_async_generator(
def is_async_callable(
func: Any,
) -> TypeGuard[Callable[..., Awaitable]]:
) -> TypeGuard[Callable[..., Awaitable[Any]]]:
"""Check if a function is async.
Args:
+32 -54
View File
@@ -122,38 +122,12 @@ def _get_annotation_description(arg_type: type) -> str | None:
return None
def _get_filtered_args(
inferred_model: type[BaseModel],
func: Callable,
*,
filter_args: Sequence[str],
include_injected: bool = True,
) -> dict:
"""Get filtered arguments from a function's signature.
Args:
inferred_model: The Pydantic model inferred from the function.
func: The function to extract arguments from.
filter_args: Arguments to exclude from the result.
include_injected: Whether to include injected arguments.
Returns:
Dictionary of filtered arguments with their schema definitions.
"""
schema = inferred_model.model_json_schema()["properties"]
valid_keys = signature(func).parameters
return {
k: schema[k]
for i, (k, param) in enumerate(valid_keys.items())
if k not in filter_args
and (i > 0 or param.name not in {"self", "cls"})
and (include_injected or not _is_injected_arg_type(param.annotation))
}
def _parse_python_function_docstring(
function: Callable, annotations: dict, *, error_on_invalid_docstring: bool = False
) -> tuple[str, dict]:
function: Callable[..., Any],
annotations: dict[str, Any],
*,
error_on_invalid_docstring: bool = False,
) -> tuple[str, dict[str, str]]:
"""Parse function and argument descriptions from a docstring.
Assumes the function docstring follows Google Python style guide.
@@ -175,7 +149,7 @@ def _parse_python_function_docstring(
def _validate_docstring_args_against_annotations(
arg_descriptions: dict, annotations: dict
arg_descriptions: dict[str, str], annotations: dict[str, Any]
) -> None:
"""Validate that docstring arguments match function annotations.
@@ -193,11 +167,11 @@ def _validate_docstring_args_against_annotations(
def _infer_arg_descriptions(
fn: Callable,
fn: Callable[..., Any],
*,
parse_docstring: bool = False,
error_on_invalid_docstring: bool = False,
) -> tuple[str, dict]:
) -> tuple[str, dict[str, str]]:
"""Infer argument descriptions from function docstring and annotations.
Args:
@@ -244,7 +218,7 @@ def _is_pydantic_annotation(annotation: Any, pydantic_version: str = "v2") -> bo
def _function_annotations_are_pydantic_v1(
signature: inspect.Signature, func: Callable
signature: inspect.Signature, func: Callable[..., Any]
) -> bool:
"""Check if all Pydantic annotations in a function are from v1.
@@ -287,7 +261,7 @@ class _SchemaConfig:
def create_schema_from_function(
model_name: str,
func: Callable,
func: Callable[..., Any],
*,
filter_args: Sequence[str] | None = None,
parse_docstring: bool = False,
@@ -415,7 +389,7 @@ content is normalized to the content of a `ToolMessage` with `status="error"`.
_EMPTY_SET: frozenset[str] = frozenset()
class BaseTool(RunnableSerializable[str | dict | ToolCall, Any]):
class BaseTool(RunnableSerializable[str | dict[str, Any] | ToolCall, Any]):
"""Base class for all LangChain tools.
This abstract class defines the interface that all LangChain tools must implement.
@@ -590,7 +564,7 @@ class ChildTool(BaseTool):
return len(keys) == 1
@property
def args(self) -> dict:
def args(self) -> dict[str, Any]:
"""Get the tool's input arguments schema.
Returns:
@@ -606,7 +580,7 @@ class ChildTool(BaseTool):
json_schema = input_schema
else:
json_schema = input_schema.model_json_schema()
return cast("dict", json_schema["properties"])
return cast("dict[str, Any]", json_schema["properties"])
@property
def tool_call_schema(self) -> ArgsSchema:
@@ -659,7 +633,7 @@ class ChildTool(BaseTool):
@override
def invoke(
self,
input: str | dict | ToolCall,
input: str | dict[str, Any] | ToolCall,
config: RunnableConfig | None = None,
**kwargs: Any,
) -> Any:
@@ -669,7 +643,7 @@ class ChildTool(BaseTool):
@override
async def ainvoke(
self,
input: str | dict | ToolCall,
input: str | dict[str, Any] | ToolCall,
config: RunnableConfig | None = None,
**kwargs: Any,
) -> Any:
@@ -679,7 +653,7 @@ class ChildTool(BaseTool):
# --- Tool ---
def _parse_input(
self, tool_input: str | dict, tool_call_id: str | None
self, tool_input: str | dict[str, Any], tool_call_id: str | None
) -> str | dict[str, Any]:
"""Parse and validate tool input using the args schema.
@@ -825,7 +799,7 @@ class ChildTool(BaseTool):
kwargs["run_manager"] = kwargs["run_manager"].get_sync()
return await run_in_executor(None, self._run, *args, **kwargs)
def _filter_injected_args(self, tool_input: dict) -> dict:
def _filter_injected_args(self, tool_input: dict[str, Any]) -> dict[str, Any]:
"""Filter out injected tool arguments from the input dictionary.
Injected arguments are those annotated with `InjectedToolArg` or its
@@ -862,8 +836,8 @@ class ChildTool(BaseTool):
return {k: v for k, v in tool_input.items() if k not in filtered_keys}
def _to_args_and_kwargs(
self, tool_input: str | dict, tool_call_id: str | None
) -> tuple[tuple, dict]:
self, tool_input: str | dict[str, Any], tool_call_id: str | None
) -> tuple[tuple[str, ...], dict[str, Any]]:
"""Convert tool input to positional and keyword arguments.
Args:
@@ -1030,7 +1004,7 @@ class ChildTool(BaseTool):
async def arun(
self,
tool_input: str | dict,
tool_input: str | dict[str, Any],
verbose: bool | None = None, # noqa: FBT001
start_color: str | None = "green",
color: str | None = "green",
@@ -1244,10 +1218,10 @@ def _handle_tool_error(
def _prep_run_args(
value: str | dict | ToolCall,
value: str | dict[str, Any] | ToolCall,
config: RunnableConfig | None,
**kwargs: Any,
) -> tuple[str | dict, dict]:
) -> tuple[str | dict[str, Any], dict[str, Any]]:
"""Prepare arguments for tool execution.
Args:
@@ -1259,12 +1233,13 @@ def _prep_run_args(
A tuple of `(tool_input, run_kwargs)`.
"""
config = ensure_config(config)
tool_input: str | dict[str, Any]
if _is_tool_call(value):
tool_call_id: str | None = cast("ToolCall", value)["id"]
tool_input: str | dict = cast("ToolCall", value)["args"].copy()
tool_input = cast("ToolCall", value)["args"].copy()
else:
tool_call_id = None
tool_input = cast("str | dict", value)
tool_input = cast("str | dict[str, Any]", value)
return (
tool_input,
dict(
@@ -1376,7 +1351,7 @@ def _stringify(content: Any) -> str:
return str(content)
def _get_type_hints(func: Callable) -> dict[str, type] | None:
def _get_type_hints(func: Callable[..., Any]) -> dict[str, type] | None:
"""Get type hints from a function, handling partial functions.
Args:
@@ -1393,7 +1368,7 @@ def _get_type_hints(func: Callable) -> dict[str, type] | None:
return None
def _get_runnable_config_param(func: Callable) -> str | None:
def _get_runnable_config_param(func: Callable[..., Any]) -> str | None:
"""Find the parameter name for `RunnableConfig` in a function.
Args:
@@ -1527,6 +1502,7 @@ def get_all_basemodel_annotations(
Returns:
`dict` of field names to their type annotations.
"""
orig_bases: tuple[type, ...]
# cls has no subscript: cls = FooBar
if isinstance(cls, type):
fields = get_fields(cls)
@@ -1540,7 +1516,7 @@ def get_all_basemodel_annotations(
continue
field_name = alias_map.get(name, name)
annotations[field_name] = param.annotation
orig_bases: tuple = getattr(cls, "__orig_bases__", ())
orig_bases = getattr(cls, "__orig_bases__", ())
# cls has subscript: cls = FooBar[int]
else:
annotations = get_all_basemodel_annotations(
@@ -1572,7 +1548,9 @@ def get_all_basemodel_annotations(
# parent_origin = class Baz,
# generic_type_vars = (type vars in Baz)
# generic_map = {type var in Baz: str}
generic_type_vars: tuple = getattr(parent_origin, "__parameters__", ())
generic_type_vars: tuple[TypeVar, ...] = getattr(
parent_origin, "__parameters__", ()
)
generic_map = dict(zip(generic_type_vars, get_args(parent), strict=False))
for field in getattr(parent_origin, "__annotations__", {}):
annotations[field] = _replace_type_vars(
+15 -13
View File
@@ -24,13 +24,13 @@ def tool(
parse_docstring: bool = False,
error_on_invalid_docstring: bool = True,
extras: dict[str, Any] | None = None,
) -> Callable[[Callable | Runnable], BaseTool]: ...
) -> Callable[[Callable[..., Any] | Runnable[Any, Any]], BaseTool]: ...
@overload
def tool(
name_or_callable: str,
runnable: Runnable,
runnable: Runnable[Any, Any],
*,
description: str | None = None,
return_direct: bool = False,
@@ -45,7 +45,7 @@ def tool(
@overload
def tool(
name_or_callable: Callable,
name_or_callable: Callable[..., Any],
*,
description: str | None = None,
return_direct: bool = False,
@@ -70,12 +70,12 @@ def tool(
parse_docstring: bool = False,
error_on_invalid_docstring: bool = True,
extras: dict[str, Any] | None = None,
) -> Callable[[Callable | Runnable], BaseTool]: ...
) -> Callable[[Callable[..., Any] | Runnable[Any, Any]], BaseTool]: ...
def tool(
name_or_callable: str | Callable | None = None,
runnable: Runnable | None = None,
name_or_callable: str | Callable[..., Any] | None = None,
runnable: Runnable[Any, Any] | None = None,
*args: Any,
description: str | None = None,
return_direct: bool = False,
@@ -85,7 +85,7 @@ def tool(
parse_docstring: bool = False,
error_on_invalid_docstring: bool = True,
extras: dict[str, Any] | None = None,
) -> BaseTool | Callable[[Callable | Runnable], BaseTool]:
) -> BaseTool | Callable[[Callable[..., Any] | Runnable[Any, Any]], BaseTool]:
"""Convert Python functions and `Runnables` to LangChain tools.
Can be used as a decorator with or without arguments to create tools from functions.
@@ -258,7 +258,7 @@ def tool(
def _create_tool_factory(
tool_name: str,
) -> Callable[[Callable | Runnable], BaseTool]:
) -> Callable[[Callable[..., Any] | Runnable[Any, Any]], BaseTool]:
"""Create a decorator that takes a callable and returns a tool.
Args:
@@ -268,7 +268,9 @@ def tool(
A function that takes a callable or `Runnable` and returns a tool.
"""
def _tool_factory(dec_func: Callable | Runnable) -> BaseTool:
def _tool_factory(
dec_func: Callable[..., Any] | Runnable[Any, Any],
) -> BaseTool:
tool_description = description
if isinstance(dec_func, Runnable):
runnable = dec_func
@@ -381,7 +383,7 @@ def tool(
# @tool(parse_docstring=True)
# def my_tool():
# pass
def _partial(func: Callable | Runnable) -> BaseTool:
def _partial(func: Callable[..., Any] | Runnable[Any, Any]) -> BaseTool:
"""Partial function that takes a `Callable` and returns a tool."""
name_ = func.get_name() if isinstance(func, Runnable) else func.__name__
tool_factory = _create_tool_factory(name_)
@@ -390,14 +392,14 @@ def tool(
return _partial
def _get_description_from_runnable(runnable: Runnable) -> str:
def _get_description_from_runnable(runnable: Runnable[Any, Any]) -> str:
"""Generate a placeholder description of a `Runnable`."""
input_schema = runnable.input_schema.model_json_schema()
return f"Takes {input_schema}."
def _get_schema_from_runnable_and_arg_types(
runnable: Runnable,
runnable: Runnable[Any, Any],
name: str,
arg_types: dict[str, type] | None = None,
) -> type[BaseModel]:
@@ -417,7 +419,7 @@ def _get_schema_from_runnable_and_arg_types(
def convert_runnable_to_tool(
runnable: Runnable,
runnable: Runnable[Any, Any],
args_schema: type[BaseModel] | None = None,
*,
name: str | None = None,
+1 -1
View File
@@ -33,7 +33,7 @@ def create_retriever_tool(
name: str,
description: str,
*,
document_prompt: BasePromptTemplate | None = None,
document_prompt: BasePromptTemplate[str] | None = None,
document_separator: str = "\n\n",
response_format: Literal["content", "content_and_artifact"] = "content",
) -> StructuredTool:
+10 -6
View File
@@ -44,7 +44,7 @@ class Tool(BaseTool):
@override
async def ainvoke(
self,
input: str | dict | ToolCall,
input: str | dict[str, Any] | ToolCall,
config: RunnableConfig | None = None,
**kwargs: Any,
) -> Any:
@@ -57,7 +57,7 @@ class Tool(BaseTool):
# --- Tool ---
@property
def args(self) -> dict:
def args(self) -> dict[str, Any]:
"""The tool's input arguments.
Returns:
@@ -70,8 +70,8 @@ class Tool(BaseTool):
return {"tool_input": {"type": "string"}}
def _to_args_and_kwargs(
self, tool_input: str | dict, tool_call_id: str | None
) -> tuple[tuple, dict]:
self, tool_input: str | dict[str, Any], tool_call_id: str | None
) -> tuple[tuple[str, ...], dict[str, Any]]:
"""Convert tool input to Pydantic model.
Args:
@@ -156,7 +156,11 @@ class Tool(BaseTool):
# TODO: this is for backwards compatibility, remove in future
def __init__(
self, name: str, func: Callable | None, description: str, **kwargs: Any
self,
name: str,
func: Callable[..., Any] | None,
description: str,
**kwargs: Any,
) -> None:
"""Initialize tool."""
super().__init__(name=name, func=func, description=description, **kwargs)
@@ -164,7 +168,7 @@ class Tool(BaseTool):
@classmethod
def from_function(
cls,
func: Callable | None,
func: Callable[..., Any] | None,
name: str, # We keep these required to support backwards compatibility
description: str,
return_direct: bool = False, # noqa: FBT001,FBT002
+3 -3
View File
@@ -59,7 +59,7 @@ class StructuredTool(BaseTool):
@override
async def ainvoke(
self,
input: str | dict | ToolCall,
input: str | dict[str, Any] | ToolCall,
config: RunnableConfig | None = None,
**kwargs: Any,
) -> Any:
@@ -132,7 +132,7 @@ class StructuredTool(BaseTool):
@classmethod
def from_function(
cls,
func: Callable | None = None,
func: Callable[..., Any] | None = None,
coroutine: Callable[..., Awaitable[Any]] | None = None,
name: str | None = None,
description: str | None = None,
@@ -263,7 +263,7 @@ class StructuredTool(BaseTool):
)
def _filter_schema_args(func: Callable) -> list[str]:
def _filter_schema_args(func: Callable[..., Any]) -> list[str]:
filter_args = list(FILTERED_ARGS)
if config_param := _get_runnable_config_param(func):
filter_args.append(config_param)
+1 -1
View File
@@ -539,7 +539,7 @@ class BaseTracer(_TracerCore, BaseCallbackHandler, ABC):
self._on_retriever_end(retrieval_run)
return retrieval_run
def __deepcopy__(self, memo: dict) -> BaseTracer:
def __deepcopy__(self, memo: dict[int, Any] | None = None) -> BaseTracer:
"""Return self."""
return self
+1 -1
View File
@@ -555,7 +555,7 @@ class _TracerCore(ABC):
retrieval_run.events.append({"name": "error", "time": retrieval_run.end_time})
return retrieval_run
def __deepcopy__(self, memo: dict) -> _TracerCore:
def __deepcopy__(self, memo: dict[int, Any] | None = None) -> _TracerCore:
"""Return self deepcopied."""
return self
@@ -56,7 +56,7 @@ class EvaluatorCallbackHandler(BaseTracer):
executor: ThreadPoolExecutor | None = None
"""The thread pool executor used for running the evaluators."""
futures: weakref.WeakSet[Future] = weakref.WeakSet()
futures: weakref.WeakSet[Future[None]] = weakref.WeakSet()
"""The set of futures representing the running evaluators."""
skip_unfinished: bool = True
@@ -98,7 +98,9 @@ def _assign_name(name: str | None, serialized: dict[str, Any] | None) -> str:
T = TypeVar("T")
class _AstreamEventsCallbackHandler(AsyncCallbackHandler, _StreamingCallbackHandler):
class _AstreamEventsCallbackHandler(
AsyncCallbackHandler, _StreamingCallbackHandler[Any]
):
"""An implementation of an async callback handler for astream events."""
def __init__(
@@ -500,7 +502,7 @@ class _AstreamEventsCallbackHandler(AsyncCallbackHandler, _StreamingCallbackHand
inputs_ = run_info.get("inputs")
generations: list[list[GenerationChunk]] | list[list[ChatGenerationChunk]]
output: dict | BaseMessage = {}
output: dict[str, Any] | BaseMessage = {}
if run_info["run_type"] == "chat_model":
generations = cast("list[list[ChatGenerationChunk]]", response.generations)
@@ -815,7 +817,9 @@ class _AstreamEventsCallbackHandler(AsyncCallbackHandler, _StreamingCallbackHand
run_info["run_type"],
)
def __deepcopy__(self, memo: dict) -> _AstreamEventsCallbackHandler:
def __deepcopy__(
self, memo: dict[int, Any] | None = None
) -> _AstreamEventsCallbackHandler:
"""Return self."""
return self
@@ -229,7 +229,7 @@ class RunLog(RunLogPatch):
T = TypeVar("T")
class LogStreamCallbackHandler(BaseTracer, _StreamingCallbackHandler):
class LogStreamCallbackHandler(BaseTracer, _StreamingCallbackHandler[Any]):
"""Tracer that streams run logs to a stream."""
def __init__(
@@ -11,14 +11,14 @@ code.
import asyncio
from asyncio import AbstractEventLoop, Queue
from collections.abc import AsyncIterator
from typing import Generic, TypeVar
from typing import Any, Generic, TypeVar
T = TypeVar("T")
class _SendStream(Generic[T]):
def __init__(
self, reader_loop: AbstractEventLoop, queue: Queue, done: object
self, reader_loop: AbstractEventLoop, queue: Queue[Any], done: object
) -> None:
"""Create a writer for the queue and done object.
@@ -84,7 +84,7 @@ class _SendStream(Generic[T]):
class _ReceiveStream(Generic[T]):
def __init__(self, queue: Queue, done: object) -> None:
def __init__(self, queue: Queue[Any], done: object) -> None:
"""Create a reader for the queue and done object.
This reader should be used in the same loop as the loop that was passed to the
@@ -126,7 +126,7 @@ class _MemoryStream(Generic[T]):
to this constructor. This will NOT be validated at run time.
"""
self._loop = loop
self._queue: asyncio.Queue = asyncio.Queue(maxsize=0)
self._queue = asyncio.Queue[Any](maxsize=0)
self._done = object()
def get_send_stream(self) -> _SendStream[T]:
+1 -1
View File
@@ -86,7 +86,7 @@ def merge_dicts(left: dict[str, Any], *others: dict[str, Any]) -> dict[str, Any]
return merged
def merge_lists(left: list | None, *others: list | None) -> list | None:
def merge_lists(left: list[Any] | None, *others: list[Any] | None) -> list[Any] | None:
"""Add many lists, handling `None`.
Args:
+1 -1
View File
@@ -277,7 +277,7 @@ class Tee(Generic[T]):
atee = Tee
class aclosing(AbstractAsyncContextManager): # noqa: N801
class aclosing(AbstractAsyncContextManager[Any]): # noqa: N801
"""Async context manager to wrap an `AsyncGenerator` that has a `aclose()` method.
Code like this:
+1 -1
View File
@@ -21,7 +21,7 @@ class StrictFormatter(Formatter):
"""
def vformat(
self, format_string: str, args: Sequence, kwargs: Mapping[str, Any]
self, format_string: str, args: Sequence[Any], kwargs: Mapping[str, Any]
) -> str:
"""Format a string using only keyword arguments.
@@ -72,7 +72,7 @@ class FunctionDescription(TypedDict):
description: str
"""A description of the function."""
parameters: dict
parameters: dict[str, Any]
"""The parameters of the function."""
@@ -86,7 +86,7 @@ class ToolDescription(TypedDict):
"""The function description."""
def _rm_titles(kv: dict, prev_key: str = "") -> dict:
def _rm_titles(kv: dict[str, Any], prev_key: str = "") -> dict[str, Any]:
"""Recursively removes `'title'` fields from a JSON schema dictionary.
Remove `'title'` fields from the input JSON schema dictionary,
@@ -121,7 +121,7 @@ def _rm_titles(kv: dict, prev_key: str = "") -> dict:
def _convert_json_schema_to_openai_function(
schema: dict,
schema: dict[str, Any],
*,
name: str | None = None,
description: str | None = None,
@@ -206,13 +206,13 @@ def _convert_pydantic_to_openai_function(
)
def _get_python_function_name(function: Callable) -> str:
def _get_python_function_name(function: Callable[..., Any]) -> str:
"""Get the name of a Python function."""
return function.__name__
def _convert_python_function_to_openai_function(
function: Callable,
function: Callable[..., Any],
) -> FunctionDescription:
"""Convert a Python function to an OpenAI function-calling API compatible dict.
@@ -243,7 +243,7 @@ def _convert_python_function_to_openai_function(
def _convert_typed_dict_to_openai_function(typed_dict: type) -> FunctionDescription:
visited: dict = {}
visited: dict[type, type] = {}
model = cast(
"type[BaseModel]",
@@ -279,7 +279,7 @@ def _convert_any_typed_dicts_to_pydantic(
description, arg_descriptions = _parse_google_docstring(
docstring, list(annotations_)
)
fields: dict = {}
fields: dict[str, Any] = {}
for arg, arg_type in annotations_.items():
if get_origin(arg_type) in {Annotated, typing_extensions.Annotated}:
annotated_args = get_args(arg_type)
@@ -373,7 +373,7 @@ def _format_tool_to_openai_function(tool: BaseTool) -> FunctionDescription:
def convert_to_openai_function(
function: Mapping[str, Any] | type | Callable | BaseTool,
function: Mapping[str, Any] | type | Callable[..., Any] | BaseTool,
*,
strict: bool | None = None,
) -> dict[str, Any]:
@@ -437,16 +437,19 @@ def convert_to_openai_function(
if function_copy and "properties" in function_copy:
oai_function["parameters"] = function_copy
elif isinstance(function, type) and is_basemodel_subclass(function):
oai_function = cast("dict", _convert_pydantic_to_openai_function(function))
oai_function = cast(
"dict[str, Any]", _convert_pydantic_to_openai_function(function)
)
elif is_typeddict(function):
oai_function = cast(
"dict", _convert_typed_dict_to_openai_function(cast("type", function))
"dict[str, Any]",
_convert_typed_dict_to_openai_function(cast("type", function)),
)
elif isinstance(function, langchain_core.tools.base.BaseTool):
oai_function = cast("dict", _format_tool_to_openai_function(function))
oai_function = cast("dict[str, Any]", _format_tool_to_openai_function(function))
elif callable(function):
oai_function = cast(
"dict", _convert_python_function_to_openai_function(function)
"dict[str, Any]", _convert_python_function_to_openai_function(function)
)
else:
if isinstance(function, dict) and (
@@ -514,7 +517,7 @@ _WellKnownOpenAITools = (
def convert_to_openai_tool(
tool: Mapping[str, Any] | type[BaseModel] | Callable | BaseTool,
tool: Mapping[str, Any] | type[BaseModel] | Callable[..., Any] | BaseTool,
*,
strict: bool | None = None,
) -> dict[str, Any]:
@@ -576,7 +579,7 @@ def convert_to_openai_tool(
def convert_to_json_schema(
schema: dict[str, Any] | type[BaseModel] | Callable | BaseTool,
schema: dict[str, Any] | type[BaseModel] | Callable[..., Any] | BaseTool,
*,
strict: bool | None = None,
) -> dict[str, Any]:
@@ -731,7 +734,7 @@ def _parse_google_docstring(
args: list[str],
*,
error_on_invalid_docstring: bool = False,
) -> tuple[str, dict]:
) -> tuple[str, dict[str, str]]:
"""Parse the function and argument descriptions from the docstring of a function.
Assumes the function docstring follows Google Python style guide.
+2 -2
View File
@@ -44,7 +44,7 @@ DEFAULT_LINK_REGEX = (
def find_all_links(
raw_html: str, *, pattern: str | re.Pattern | None = None
raw_html: str, *, pattern: str | re.Pattern[str] | None = None
) -> list[str]:
"""Extract all links from a raw HTML string.
@@ -64,7 +64,7 @@ def extract_sub_links(
url: str,
*,
base_url: str | None = None,
pattern: str | re.Pattern | None = None,
pattern: str | re.Pattern[str] | None = None,
prevent_outside: bool = True,
exclude_prefixes: Sequence[str] = (),
continue_on_failure: bool = False,
+1 -1
View File
@@ -12,7 +12,7 @@ _TEXT_COLOR_MAPPING = {
def get_color_mapping(
items: list[str], excluded_colors: list | None = None
items: list[str], excluded_colors: list[str] | None = None
) -> dict[str, str]:
"""Get mapping for items to a support color.
+3 -1
View File
@@ -191,7 +191,9 @@ def _parse_json(
return parser(json_str)
def parse_and_check_json_markdown(text: str, expected_keys: list[str]) -> dict:
def parse_and_check_json_markdown(
text: str, expected_keys: list[str]
) -> dict[str, Any]:
"""Parse and check a JSON string from a Markdown string.
Checks that it contains the expected keys.
@@ -9,7 +9,7 @@ if TYPE_CHECKING:
from collections.abc import Sequence
def _retrieve_ref(path: str, schema: dict) -> list | dict:
def _retrieve_ref(path: str, schema: dict[str, Any]) -> list[Any] | dict[Any, Any]:
"""Retrieve a referenced object from a JSON schema using a path.
Resolves JSON schema references (e.g., `'#/definitions/MyType'`) by traversing the
@@ -33,7 +33,7 @@ def _retrieve_ref(path: str, schema: dict) -> list | dict:
"with #."
)
raise ValueError(msg)
out: list | dict = schema
out: list[Any] | dict[Any, Any] = schema
for component in components[1:]:
if component in out:
if isinstance(out, list):
@@ -186,11 +186,11 @@ def _dereference_refs_helper(
def dereference_refs(
schema_obj: dict,
schema_obj: dict[str, Any],
*,
full_schema: dict | None = None,
full_schema: dict[str, Any] | None = None,
skip_keys: Sequence[str] | None = None,
) -> dict:
) -> dict[str, Any]:
"""Resolve and inline JSON Schema `$ref` references in a schema object.
This function processes a JSON Schema and resolves all `$ref` references by
@@ -266,7 +266,7 @@ def dereference_refs(
keys_to_skip = list(skip_keys) if skip_keys is not None else ["$defs"]
shallow = skip_keys is None
return cast(
"dict",
"dict[str, Any]",
_dereference_refs_helper(
schema_obj, full, None, keys_to_skip, shallow_refs=shallow
),
+7 -5
View File
@@ -126,7 +126,9 @@ def is_basemodel_instance(obj: Any) -> bool:
# How to type hint this?
def pre_init(func: Callable) -> Any:
def pre_init(
func: Callable[[Any, dict[str, Any]], Any],
) -> Callable[[Any, dict[str, Any]], Any]:
"""Decorator to run a function before model initialization.
Args:
@@ -202,9 +204,9 @@ class _IgnoreUnserializable(GenerateJsonSchema):
def _create_subset_model_v1(
name: str,
model: type[BaseModelV1],
field_names: list,
field_names: list[str],
*,
descriptions: dict | None = None,
descriptions: dict[str, str] | None = None,
fn_description: str | None = None,
) -> type[BaseModelV1]:
"""Create a Pydantic model with only a subset of model's fields."""
@@ -233,7 +235,7 @@ def _create_subset_model_v2(
model: type[BaseModel],
field_names: list[str],
*,
descriptions: dict | None = None,
descriptions: dict[str, str] | None = None,
fn_description: str | None = None,
) -> type[BaseModel]:
"""Create a Pydantic model with a subset of the model fields."""
@@ -283,7 +285,7 @@ def _create_subset_model(
model: TypeBaseModel,
field_names: list[str],
*,
descriptions: dict | None = None,
descriptions: dict[str, str] | None = None,
fn_description: str | None = None,
) -> type[BaseModel]:
"""Create subset model using the same pydantic version as the input model.
+1 -1
View File
@@ -22,7 +22,7 @@ def stringify_value(val: Any) -> str:
return str(val)
def stringify_dict(data: dict) -> str:
def stringify_dict(data: dict[Any, Any]) -> str:
"""Stringify a dictionary.
Args:
+5 -4
View File
@@ -1,17 +1,18 @@
"""Usage utilities."""
from collections.abc import Callable
from typing import Any
def _dict_int_op(
left: dict,
right: dict,
left: dict[str, Any],
right: dict[str, Any],
op: Callable[[int, int], int],
*,
default: int = 0,
depth: int = 0,
max_depth: int = 100,
) -> dict:
) -> dict[str, Any]:
"""Apply an integer operation to corresponding values in two dictionaries.
Recursively combines two dictionaries by applying the given operation to integer
@@ -36,7 +37,7 @@ def _dict_int_op(
if depth >= max_depth:
msg = f"{max_depth=} exceeded, unable to combine dicts."
raise ValueError(msg)
combined: dict = {}
combined: dict[str, Any] = {}
for k in set(left).union(right):
if isinstance(left.get(k, default), int) and isinstance(
right.get(k, default), int
+2 -2
View File
@@ -21,7 +21,7 @@ from langchain_core.utils.pydantic import (
)
def xor_args(*arg_groups: tuple[str, ...]) -> Callable:
def xor_args(*arg_groups: tuple[str, ...]) -> Callable[..., Any]:
"""Validate specified keyword args are mutually exclusive.
Args:
@@ -31,7 +31,7 @@ def xor_args(*arg_groups: tuple[str, ...]) -> Callable:
Decorator that validates the specified keyword args are mutually exclusive.
"""
def decorator(func: Callable) -> Callable:
def decorator(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(*args: Any, **kwargs: Any) -> Any:
"""Validate exactly one arg in each group is not None."""
@@ -46,7 +46,7 @@ class VectorStore(ABC):
def add_texts(
self,
texts: Iterable[str],
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
*,
ids: list[str] | None = None,
**kwargs: Any,
@@ -185,7 +185,7 @@ class VectorStore(ABC):
async def aadd_texts(
self,
texts: Iterable[str],
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
*,
ids: list[str] | None = None,
**kwargs: Any,
@@ -849,7 +849,7 @@ class VectorStore(ABC):
cls: type[VST],
texts: list[str],
embedding: Embeddings,
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
*,
ids: list[str] | None = None,
**kwargs: Any,
@@ -872,7 +872,7 @@ class VectorStore(ABC):
cls,
texts: list[str],
embedding: Embeddings,
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
*,
ids: list[str] | None = None,
**kwargs: Any,
@@ -970,7 +970,7 @@ class VectorStoreRetriever(BaseRetriever):
search_type: str = "similarity"
"""Type of search to perform."""
search_kwargs: dict = Field(default_factory=dict)
search_kwargs: dict[str, Any] = Field(default_factory=dict)
"""Keyword arguments to pass to the search function."""
allowed_search_types: ClassVar[Collection[str]] = (
@@ -985,7 +985,7 @@ class VectorStoreRetriever(BaseRetriever):
@model_validator(mode="before")
@classmethod
def validate_search_type(cls, values: dict) -> Any:
def validate_search_type(cls, values: dict[str, Any]) -> Any:
"""Validate search type.
Args:
@@ -489,7 +489,7 @@ class InMemoryVectorStore(VectorStore):
cls,
texts: list[str],
embedding: Embeddings,
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
**kwargs: Any,
) -> InMemoryVectorStore:
store = cls(
@@ -504,7 +504,7 @@ class InMemoryVectorStore(VectorStore):
cls,
texts: list[str],
embedding: Embeddings,
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
**kwargs: Any,
) -> InMemoryVectorStore:
store = cls(
+13 -7
View File
@@ -10,7 +10,7 @@ from __future__ import annotations
import logging
import warnings
from typing import TYPE_CHECKING, cast
from typing import TYPE_CHECKING
try:
import numpy as np
@@ -27,12 +27,16 @@ except ImportError:
_HAS_SIMSIMD = False
if TYPE_CHECKING:
Matrix = list[list[float]] | list[np.ndarray] | np.ndarray
import numpy.typing as npt
Matrix = (
list[list[float]] | list[npt.NDArray[np.floating]] | npt.NDArray[np.floating]
)
logger = logging.getLogger(__name__)
def _cosine_similarity(x: Matrix, y: Matrix) -> np.ndarray:
def _cosine_similarity(x: Matrix, y: Matrix) -> npt.NDArray[np.floating]:
"""Row-wise cosine similarity between two equal-width matrices.
Args:
@@ -91,12 +95,14 @@ def _cosine_similarity(x: Matrix, y: Matrix) -> np.ndarray:
y_norm = np.linalg.norm(y, axis=1)
# Ignore divide by zero errors run time warnings as those are handled below.
with np.errstate(divide="ignore", invalid="ignore"):
similarity = np.dot(x, y.T) / np.outer(x_norm, y_norm)
similarity: npt.NDArray[np.floating] = np.dot(x, y.T) / np.outer(
x_norm, y_norm
)
if np.isnan(similarity).all():
msg = "NaN values found, please remove the NaN values and try again"
raise ValueError(msg) from None
similarity[np.isnan(similarity) | np.isinf(similarity)] = 0.0
return cast("np.ndarray", similarity)
return similarity
x = np.array(x, dtype=np.float32)
y = np.array(y, dtype=np.float32)
@@ -104,8 +110,8 @@ def _cosine_similarity(x: Matrix, y: Matrix) -> np.ndarray:
def maximal_marginal_relevance(
query_embedding: np.ndarray,
embedding_list: list,
query_embedding: npt.NDArray[np.floating],
embedding_list: list[list[float]],
lambda_mult: float = 0.5,
k: int = 4,
) -> list[int]:
@@ -1,4 +1,7 @@
from pathlib import Path
from langchain_core.documents import Document
from langchain_core.documents.base import Blob
def test_init() -> None:
@@ -10,3 +13,15 @@ def test_init() -> None:
Document(page_content="foo", id=1),
]:
assert isinstance(doc, Document)
def test_metadata_allows_non_string_keys(tmp_path: Path) -> None:
metadata = {1: "one"}
doc = Document(page_content="foo", metadata=metadata)
blob_from_data = Blob.from_data("foo", metadata=metadata)
blob_from_path = Blob.from_path(tmp_path / "foo.txt", metadata=metadata)
assert doc.metadata == metadata
assert blob_from_data.metadata == metadata
assert blob_from_path.metadata == metadata
@@ -15,7 +15,7 @@ from langchain_core.vectorstores import VectorStore
class DummyVectorStore(VectorStore):
def __init__(self, init_arg: str | None = None):
self.texts: list[str] = []
self.metadatas: list[dict] = []
self.metadatas: list[dict[str, Any]] = []
self._embeddings: Embeddings | None = None
self.init_arg = init_arg
@@ -27,7 +27,7 @@ class DummyVectorStore(VectorStore):
def add_texts(
self,
texts: Iterable[str],
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
**kwargs: Any,
) -> list[str]:
self.texts.extend(texts)
@@ -66,7 +66,7 @@ class DummyVectorStore(VectorStore):
cls,
texts: list[str],
embedding: Embeddings,
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
**kwargs: Any,
) -> "DummyVectorStore":
store = DummyVectorStore(**kwargs)
+15 -7
View File
@@ -1,4 +1,4 @@
from typing import cast
from typing import Any, cast
from langchain_core.load import dumpd, load
from langchain_core.messages import AIMessage, AIMessageChunk
@@ -355,11 +355,11 @@ def test_content_blocks() -> None:
assert chunk.content_blocks == chunk.tool_calls
# test v1 content
chunk_1.content = cast("str | list[str | dict]", chunk_1.content_blocks)
chunk_1.content = cast("list[str | dict[str, Any]]", chunk_1.content_blocks)
assert len(chunk_1.content) == 1
chunk_1.content[0]["extras"] = {"baz": "qux"} # type: ignore[index]
chunk_1.response_metadata["output_version"] = "v1"
chunk_2.content = cast("str | list[str | dict]", chunk_2.content_blocks)
chunk_2.content = cast("list[str | dict[str, Any]]", chunk_2.content_blocks)
chunk = chunk_1 + chunk_2 + chunk_3
assert chunk.content == [
@@ -379,8 +379,12 @@ def test_content_blocks() -> None:
standard_content_2: list[types.ContentBlock] = [
{"type": "non_standard", "index": 0, "value": {"foo": "baz"}}
]
chunk_1 = AIMessageChunk(content=cast("str | list[str | dict]", standard_content_1))
chunk_2 = AIMessageChunk(content=cast("str | list[str | dict]", standard_content_2))
chunk_1 = AIMessageChunk(
content=cast("list[str | dict[str, Any]]", standard_content_1)
)
chunk_2 = AIMessageChunk(
content=cast("list[str | dict[str, Any]]", standard_content_2)
)
merged_chunk = chunk_1 + chunk_2
assert merged_chunk.content == [
{"type": "non_standard", "index": 0, "value": {"foo": "bar baz"}},
@@ -467,8 +471,12 @@ def test_content_blocks() -> None:
}
]
standard_content_2 = [{"type": "non_standard", "value": {"foo": "bar"}, "index": 0}]
chunk_1 = AIMessageChunk(content=cast("str | list[str | dict]", standard_content_1))
chunk_2 = AIMessageChunk(content=cast("str | list[str | dict]", standard_content_2))
chunk_1 = AIMessageChunk(
content=cast("list[str | dict[str, Any]]", standard_content_1)
)
chunk_2 = AIMessageChunk(
content=cast("list[str | dict[str, Any]]", standard_content_2)
)
merged_chunk = chunk_1 + chunk_2
assert merged_chunk.content == [
{
@@ -758,7 +758,8 @@ class FakeTokenCountingModel(FakeChatModel):
def get_num_tokens_from_messages(
self,
messages: list[BaseMessage],
tools: Sequence[dict[str, Any] | type | Callable | BaseTool] | None = None,
tools: Sequence[dict[str, Any] | type | Callable[..., Any] | BaseTool]
| None = None,
) -> int:
return dummy_token_counter(messages)
@@ -1288,7 +1289,7 @@ def test_convert_to_openai_messages_invalid_block() -> None:
def test_handle_openai_responses_blocks() -> None:
blocks: str | list[str | dict] = [
blocks: str | list[str | dict[str, Any]] = [
{"type": "reasoning", "id": "1"},
{
"type": "function_call",
@@ -6,7 +6,7 @@ import pytest
from pydantic import BaseModel
from typing_extensions import override
from langchain_core.language_models import FakeListChatModel
from langchain_core.language_models import FakeListChatModel, LanguageModelInput
from langchain_core.load.dump import dumps
from langchain_core.load.load import loads
from langchain_core.messages import HumanMessage
@@ -29,8 +29,8 @@ class FakeStructuredChatModel(FakeListChatModel):
@override
def with_structured_output(
self, schema: dict | type[BaseModel], **kwargs: Any
) -> Runnable:
self, schema: dict[str, Any] | type[BaseModel], **kwargs: Any
) -> Runnable[LanguageModelInput, dict[str, Any] | BaseModel]:
return RunnableLambda(partial(_fake_runnable, schema=schema, **kwargs))
@property
+2 -2
View File
@@ -73,7 +73,7 @@ def _remove_enum(obj: Any) -> None:
_remove_enum(item)
def _schema(obj: Any) -> dict:
def _schema(obj: Any) -> dict[str, Any]:
"""Return the schema of the object."""
# Remap to old style schema
if isclass(obj):
@@ -99,7 +99,7 @@ def _schema(obj: Any) -> dict:
raise TypeError(msg)
def _remove_additionalproperties(schema: dict) -> dict[str, Any]:
def _remove_additionalproperties(schema: dict[str, Any]) -> dict[str, Any]:
"""Remove `"additionalProperties": True` from dicts in the schema.
Pydantic 2.11 and later versions include `"additionalProperties": True` when
@@ -335,14 +335,16 @@ class FakeStructuredOutputModel(BaseChatModel):
@override
def bind_tools(
self,
tools: Sequence[dict[str, Any] | type[BaseModel] | Callable | BaseTool],
tools: Sequence[
dict[str, Any] | type[BaseModel] | Callable[..., Any] | BaseTool
],
**kwargs: Any,
) -> Runnable[LanguageModelInput, AIMessage]:
return self.bind(tools=tools)
@override
def with_structured_output(
self, schema: dict | type[BaseModel], **kwargs: Any
self, schema: dict[str, Any] | type[BaseModel], **kwargs: Any
) -> Runnable[LanguageModelInput, dict[str, int] | BaseModel]:
return RunnableLambda(lambda _: {"foo": self.foo})
@@ -368,7 +370,9 @@ class FakeModel(BaseChatModel):
@override
def bind_tools(
self,
tools: Sequence[dict[str, Any] | type[BaseModel] | Callable | BaseTool],
tools: Sequence[
dict[str, Any] | type[BaseModel] | Callable[..., Any] | BaseTool
],
**kwargs: Any,
) -> Runnable[LanguageModelInput, AIMessage]:
return self.bind(tools=tools)
@@ -226,7 +226,9 @@ def test_graph_sequence_map(snapshot: SnapshotAssertion) -> None:
str_parser = StrOutputParser()
xml_parser = XMLOutputParser()
def conditional_str_parser(value: str) -> Runnable[BaseMessage | str, str]:
def conditional_str_parser(
value: str,
) -> Runnable[BaseMessage | str, str | dict[str, Any]]:
if value == "a":
return str_parser
return xml_parser
@@ -3770,6 +3770,40 @@ async def test_deep_astream_assign() -> None:
}
def _empty_mapper_assign() -> RunnableAssign:
"""Build an assign whose mapper yields zero chunks.
The map output stream is started with `next(map_output, None)` /
`anext(map_output, None)`, so `None` is the exhaustion sentinel rather than
a real chunk. Both `stream` and `astream` must guard against yielding that
sentinel into the output stream.
"""
def empty_gen(it: Iterator[Any]) -> Iterator[dict[str, Any]]:
for _ in it:
pass
yield from () # consume input, yield nothing
async def aempty_gen(it: AsyncIterator[Any]) -> AsyncIterator[dict[str, Any]]:
async for _ in it:
pass
return
yield # pragma: no cover # make this an async generator function
return RunnablePassthrough.assign(foo=RunnableGenerator(empty_gen, aempty_gen))
def test_stream_assign_empty_mapper() -> None:
"""An assign whose mapper yields no chunks must not emit `None` (sync)."""
assert list(_empty_mapper_assign().stream({"a": 1})) == [{"a": 1}]
async def test_astream_assign_empty_mapper() -> None:
"""An assign whose mapper yields no chunks must not emit `None` (async)."""
chunks = [chunk async for chunk in _empty_mapper_assign().astream({"a": 1})]
assert chunks == [{"a": 1}]
def test_runnable_sequence_transform() -> None:
llm = FakeStreamingListLLM(responses=["foo-lish"])
@@ -5374,7 +5408,7 @@ async def test_ainvoke_on_returned_runnable() -> None:
def test_invoke_stream_passthrough_assign_trace() -> None:
def idchain_sync(_input: dict, /) -> bool:
def idchain_sync(_input: dict[str, Any], /) -> bool:
return False
chain = RunnablePassthrough.assign(urls=idchain_sync)
@@ -5394,7 +5428,7 @@ def test_invoke_stream_passthrough_assign_trace() -> None:
async def test_ainvoke_astream_passthrough_assign_trace() -> None:
def idchain_sync(_input: dict, /) -> bool:
def idchain_sync(_input: dict[str, Any], /) -> bool:
return False
chain = RunnablePassthrough.assign(urls=idchain_sync)
@@ -78,7 +78,7 @@ def _assert_events_equal_allow_superset_metadata(
async def test_event_stream_with_simple_function_tool() -> None:
"""Test the event stream with a function and tool."""
def foo(_: int) -> dict:
def foo(_: int) -> dict[str, Any]:
"""Foo."""
return {"x": 5}
@@ -1076,12 +1076,14 @@ async def test_event_streaming_with_tools() -> None:
return "world"
@tool
def with_parameters(x: int, y: str) -> dict:
def with_parameters(x: int, y: str) -> dict[str, Any]:
"""A tool that does nothing."""
return {"x": x, "y": y}
@tool
def with_parameters_and_callbacks(x: int, y: str, callbacks: Callbacks) -> dict:
def with_parameters_and_callbacks(
x: int, y: str, callbacks: Callbacks
) -> dict[str, Any]:
"""A tool that does nothing."""
_ = callbacks
return {"x": x, "y": y}
@@ -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:
def foo(x: int) -> dict[str, Any]:
"""Foo."""
_ = x
return {"x": 5}
@@ -1097,12 +1097,14 @@ async def test_event_streaming_with_tools() -> None:
return "world"
@tool
def with_parameters(x: int, y: str) -> dict:
def with_parameters(x: int, y: str) -> dict[str, Any]:
"""A tool that does nothing."""
return {"x": x, "y": y}
@tool
def with_parameters_and_callbacks(x: int, y: str, callbacks: Callbacks) -> dict:
def with_parameters_and_callbacks(
x: int, y: str, callbacks: Callbacks
) -> dict[str, Any]:
"""A tool that does nothing."""
_ = callbacks
return {"x": x, "y": y}
@@ -421,7 +421,8 @@ class TestRunnableSequenceParallelTraceNesting:
ids=["invoke", "stream", "batch"],
)
def test_sync(
self, method: Callable[[RunnableLambda, list[BaseCallbackHandler]], int]
self,
method: Callable[[RunnableLambda[int, int], list[BaseCallbackHandler]], int],
) -> None:
def other_thing(_: int) -> Generator[int, None, None]:
yield 1
@@ -458,7 +459,8 @@ class TestRunnableSequenceParallelTraceNesting:
async def test_async(
self,
method: Callable[
[RunnableLambda, list[BaseCallbackHandler]], Coroutine[Any, Any, int]
[RunnableLambda[int, int], list[BaseCallbackHandler]],
Coroutine[Any, Any, int],
],
) -> None:
async def other_thing(_: int) -> AsyncGenerator[int, None]:
@@ -56,7 +56,7 @@ def test_nonlocals() -> None:
def my_func4(value: str) -> str:
return global_agent.invoke(value)
def my_func5() -> tuple[Callable[[str], str], RunnableLambda]:
def my_func5() -> tuple[Callable[[str], str], RunnableLambda[str, str]]:
global_agent = RunnableLambda[str, str](lambda x: x * 3)
def my_func6(value: str) -> str:
@@ -23,7 +23,7 @@ class TestSyncInMemoryStore(BaseStoreSyncTests[Any]):
return "value1", "value2", "value3"
class TestAsyncInMemoryStore(BaseStoreAsyncTests):
class TestAsyncInMemoryStore(BaseStoreAsyncTests[Any]):
@pytest.fixture
@override
async def kv_store(self) -> InMemoryStore:
+7 -3
View File
@@ -1,5 +1,5 @@
import uuid
from typing import get_args
from typing import Any, get_args
import pytest
@@ -488,7 +488,7 @@ def test_message_chunk_to_message() -> None:
def test_tool_calls_merge() -> None:
chunks: list[dict] = [
chunks: list[dict[str, Any]] = [
{"content": ""},
{
"content": "",
@@ -1092,7 +1092,11 @@ def test_tool_message_str() -> None:
),
],
)
def test_merge_content(first: list | str, others: list, expected: list | str) -> None:
def test_merge_content(
first: str | list[str | dict[str, Any]],
others: str | list[str | dict[str, Any]],
expected: str | list[str | dict[str, Any]],
) -> None:
actual = merge_content(first, *others)
assert actual == expected
@@ -12,7 +12,7 @@ from langchain_core.prompt_values import ChatPromptValueConcrete
def test_chat_prompt_value_concrete() -> None:
messages: list = [
messages = [
AIMessage("foo"),
HumanMessage("foo"),
SystemMessage("foo"),
+44 -28
View File
@@ -113,7 +113,7 @@ class _MockSchema(BaseModel):
arg1: int
arg2: bool
arg3: dict | None = None
arg3: dict[str, Any] | None = None
class _MockStructuredTool(BaseTool):
@@ -122,10 +122,12 @@ class _MockStructuredTool(BaseTool):
description: str = "A Structured Tool"
@override
def _run(self, *, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
def _run(self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None) -> str:
return f"{arg1} {arg2} {arg3}"
async def _arun(self, *, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
async def _arun(
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
) -> str:
raise NotImplementedError
@@ -166,11 +168,13 @@ def test_misannotated_base_tool_raises_error() -> None:
description: str = "A Structured Tool"
@override
def _run(self, *, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
def _run(
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
) -> str:
return f"{arg1} {arg2} {arg3}"
async def _arun(
self, *, arg1: int, arg2: bool, arg3: dict | None = None
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
) -> str:
raise NotImplementedError
@@ -184,11 +188,13 @@ def test_forward_ref_annotated_base_tool_accepted() -> None:
description: str = "A Structured Tool"
@override
def _run(self, *, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
def _run(
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
) -> str:
return f"{arg1} {arg2} {arg3}"
async def _arun(
self, *, arg1: int, arg2: bool, arg3: dict | None = None
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
) -> str:
raise NotImplementedError
@@ -202,11 +208,13 @@ def test_subclass_annotated_base_tool_accepted() -> None:
description: str = "A Structured Tool"
@override
def _run(self, *, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
def _run(
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
) -> str:
return f"{arg1} {arg2} {arg3}"
async def _arun(
self, *, arg1: int, arg2: bool, arg3: dict | None = None
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
) -> str:
raise NotImplementedError
@@ -219,7 +227,7 @@ def test_decorator_with_specified_schema() -> None:
"""Test that manually specified schemata are passed through to the tool."""
@tool(args_schema=_MockSchema)
def tool_func(*, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
def tool_func(*, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None) -> str:
return f"{arg1} {arg2} {arg3}"
assert isinstance(tool_func, BaseTool)
@@ -238,10 +246,12 @@ def test_decorator_with_specified_schema_pydantic_v1() -> None:
arg1: int
arg2: bool
arg3: dict | None = None
arg3: dict[str, Any] | None = None
@tool(args_schema=cast("ArgsSchema", _MockSchemaV1))
def tool_func_v1(*, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
def tool_func_v1(
*, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
) -> str:
return f"{arg1} {arg2} {arg3}"
assert isinstance(tool_func_v1, BaseTool)
@@ -253,7 +263,7 @@ def test_decorated_function_schema_equivalent() -> None:
@tool
def structured_tool_input(
*, arg1: int, arg2: bool, arg3: dict | None = None
*, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
) -> str:
"""Return the arguments directly."""
return f"{arg1} {arg2} {arg3}"
@@ -322,7 +332,7 @@ def test_structured_args_decorator_no_infer_schema() -> None:
@tool(infer_schema=False)
def structured_tool_input(
arg1: int, arg2: float | datetime, opt_arg: dict | None = None
arg1: int, arg2: float | datetime, opt_arg: dict[str, Any] | None = None
) -> str:
"""Return the arguments directly."""
return f"{arg1}, {arg2}, {opt_arg}"
@@ -362,7 +372,7 @@ def test_structured_tool_types_parsed() -> None:
def structured_tool(
some_enum: SomeEnum,
some_base_model: SomeBaseModel,
) -> dict:
) -> dict[str, Any]:
"""Return the arguments directly."""
return {
"some_enum": some_enum,
@@ -1032,7 +1042,7 @@ def test_validation_error_handling_non_validation_error(
def _parse_input(
self,
tool_input: str | dict,
tool_input: str | dict[str, Any],
tool_call_id: str | None,
) -> str | dict[str, Any]:
raise NotImplementedError
@@ -1098,7 +1108,7 @@ async def test_async_validation_error_handling_non_validation_error(
def _parse_input(
self,
tool_input: str | dict,
tool_input: str | dict[str, Any],
tool_call_id: str | None,
) -> str | dict[str, Any]:
raise NotImplementedError
@@ -1145,9 +1155,11 @@ def test_optional_subset_model_rewrite() -> None:
({"bar": "bar", "baz": None}, {"bar": "bar", "baz": None, "buzz": "buzz"}),
],
)
def test_tool_invoke_optional_args(inputs: dict, expected: dict | None) -> None:
def test_tool_invoke_optional_args(
inputs: dict[str, Any], expected: dict[str, Any] | None
) -> None:
@tool
def foo(bar: str, baz: int | None = 3, buzz: str | None = "buzz") -> dict:
def foo(bar: str, baz: int | None = 3, buzz: str | None = "buzz") -> dict[str, Any]:
"""The foo."""
return {
"bar": bar,
@@ -2283,9 +2295,13 @@ def test__get_all_basemodel_annotations_v2(*, use_v1_namespace: bool) -> None:
return "foo"
class ModelC(Mixin, ModelB):
c: dict
c: dict[str, Any]
expected = {"a": str, "b": Annotated[ModelA[dict[str, Any]], "foo"], "c": dict}
expected = {
"a": str,
"b": Annotated[ModelA[dict[str, Any]], "foo"],
"c": dict[str, Any],
}
actual = get_all_basemodel_annotations(ModelC)
assert actual == expected
@@ -2309,7 +2325,7 @@ def test__get_all_basemodel_annotations_v2(*, use_v1_namespace: bool) -> None:
expected = {
"a": str,
"b": Annotated[ModelA[dict[str, Any]], "foo"],
"c": dict,
"c": dict[str, Any],
"d": str | int | None,
}
actual = get_all_basemodel_annotations(ModelD)
@@ -2318,7 +2334,7 @@ def test__get_all_basemodel_annotations_v2(*, use_v1_namespace: bool) -> None:
expected = {
"a": str,
"b": Annotated[ModelA[dict[str, Any]], "foo"],
"c": dict,
"c": dict[str, Any],
"d": int | None,
}
actual = get_all_basemodel_annotations(ModelD[int])
@@ -2342,7 +2358,7 @@ def test_tool_annotations_preserved() -> None:
"""Test that annotations are preserved when creating a tool."""
@tool
def my_tool(val: int, other_val: Annotated[dict, "my annotation"]) -> str:
def my_tool(val: int, other_val: Annotated[dict[str, Any], "my annotation"]) -> str:
"""Tool docstring."""
return "foo"
@@ -3080,7 +3096,7 @@ def test_tool_args_schema_with_annotated_type() -> None:
class CallbackHandlerWithInputCapture(FakeCallbackHandler):
"""Callback handler that captures inputs passed to on_tool_start."""
captured_inputs: list[dict | None] = Field(default_factory=list)
captured_inputs: list[dict[str, Any] | None] = Field(default_factory=list)
def on_tool_start(
self,
@@ -3114,7 +3130,7 @@ def test_filter_injected_args_from_callbacks() -> None:
@tool
def search_tool(
query: str,
state: Annotated[dict, InjectedToolArg()],
state: Annotated[dict[str, Any], InjectedToolArg()],
) -> str:
"""Search with injected state.
@@ -3182,7 +3198,7 @@ def test_filter_multiple_injected_args() -> None:
def complex_tool(
query: str,
limit: int,
state: Annotated[dict, InjectedToolArg()],
state: Annotated[dict[str, Any], InjectedToolArg()],
context: Annotated[str, InjectedToolArg()],
run_manager: CallbackManagerForToolRun | None = None,
) -> str:
@@ -3250,7 +3266,7 @@ async def test_filter_injected_args_async() -> None:
@tool
async def async_search_tool(
query: str,
state: Annotated[dict, InjectedToolArg()],
state: Annotated[dict[str, Any], InjectedToolArg()],
) -> str:
"""Async search with injected state.
@@ -12,7 +12,7 @@ async def test_same_event_loop() -> None:
This is the easy case.
"""
reader_loop = asyncio.get_event_loop()
channel = _MemoryStream[dict](reader_loop)
channel = _MemoryStream[dict[str, int | float]](reader_loop)
writer = channel.get_send_stream()
reader = channel.get_receive_stream()
@@ -30,7 +30,7 @@ async def test_same_event_loop() -> None:
)
await writer.aclose()
async def consumer() -> AsyncIterator[dict]:
async def consumer() -> AsyncIterator[dict[str, int | float]]:
tic = time.time()
async for item in reader:
toc = time.time()
@@ -63,7 +63,7 @@ async def test_same_event_loop() -> None:
async def test_queue_for_streaming_via_sync_call() -> None:
"""Test via async -> sync -> async path."""
reader_loop = asyncio.get_event_loop()
channel = _MemoryStream[dict](reader_loop)
channel = _MemoryStream[dict[str, int | float]](reader_loop)
writer = channel.get_send_stream()
reader = channel.get_receive_stream()
@@ -85,7 +85,7 @@ async def test_queue_for_streaming_via_sync_call() -> None:
"""Blocking sync call."""
asyncio.run(producer())
async def consumer() -> AsyncIterator[dict]:
async def consumer() -> AsyncIterator[dict[str, int | float]]:
tic = time.time()
async for item in reader:
toc = time.time()
@@ -26,7 +26,7 @@ from pydantic import BaseModel, ConfigDict, Field
from pydantic.errors import PydanticInvalidForJsonSchema
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from langchain_core.runnables import Runnable, RunnableLambda
from langchain_core.runnables import RunnableLambda
from langchain_core.tools import BaseTool, StructuredTool, Tool, tool
from langchain_core.utils.function_calling import (
_convert_typed_dict_to_openai_function,
@@ -49,7 +49,7 @@ def pydantic() -> type[BaseModel]:
@pytest.fixture
def annotated_function() -> Callable:
def annotated_function() -> Callable[[int, Literal["bar", "baz"]], None]:
def dummy_function(
arg1: ExtensionsAnnotated[int, "foo"],
arg2: ExtensionsAnnotated[Literal["bar", "baz"], "one of 'bar', 'baz'"],
@@ -60,7 +60,7 @@ def annotated_function() -> Callable:
@pytest.fixture
def function() -> Callable:
def function() -> Callable[[int, Literal["bar", "baz"]], None]:
def dummy_function(arg1: int, arg2: Literal["bar", "baz"]) -> None:
"""Dummy function.
@@ -73,7 +73,7 @@ def function() -> Callable:
@pytest.fixture
def function_docstring_annotations() -> Callable:
def function_docstring_annotations() -> Callable[[int, Literal["bar", "baz"]], None]:
def dummy_function(arg1: int, arg2: Literal["bar", "baz"]) -> None:
"""Dummy function.
@@ -85,13 +85,14 @@ def function_docstring_annotations() -> Callable:
return dummy_function
@pytest.fixture
def runnable() -> Runnable:
class Args(ExtensionsTypedDict):
arg1: ExtensionsAnnotated[int, "foo"]
arg2: ExtensionsAnnotated[Literal["bar", "baz"], "one of 'bar', 'baz'"]
class _Args(ExtensionsTypedDict):
arg1: ExtensionsAnnotated[int, "foo"]
arg2: ExtensionsAnnotated[Literal["bar", "baz"], "one of 'bar', 'baz'"]
def dummy_function(input_dict: Args) -> None:
@pytest.fixture
def runnable() -> RunnableLambda[_Args, None]:
def dummy_function(input_dict: _Args) -> None:
pass
return RunnableLambda(dummy_function)
@@ -229,7 +230,7 @@ def dummy_extensions_typed_dict_docstring() -> type:
@pytest.fixture
def json_schema() -> dict:
def json_schema() -> dict[str, Any]:
return {
"title": "dummy_function",
"description": "Dummy function.",
@@ -247,7 +248,7 @@ def json_schema() -> dict:
@pytest.fixture
def anthropic_tool() -> dict:
def anthropic_tool() -> dict[str, Any]:
return {
"name": "dummy_function",
"description": "Dummy function.",
@@ -267,7 +268,7 @@ def anthropic_tool() -> dict:
@pytest.fixture
def bedrock_converse_tool() -> dict:
def bedrock_converse_tool() -> dict[str, Any]:
return {
"toolSpec": {
"name": "dummy_function",
@@ -313,17 +314,17 @@ class DummyWithClassMethod:
def test_convert_to_openai_function(
pydantic: type[BaseModel],
function: Callable,
function_docstring_annotations: Callable,
function: Callable[[int, Literal["bar", "baz"]], None],
function_docstring_annotations: Callable[[int, Literal["bar", "baz"]], None],
dummy_structured_tool: StructuredTool,
dummy_structured_tool_args_schema_dict: StructuredTool,
dummy_tool: BaseTool,
json_schema: dict,
anthropic_tool: dict,
bedrock_converse_tool: dict,
annotated_function: Callable,
json_schema: dict[str, Any],
anthropic_tool: dict[str, Any],
bedrock_converse_tool: dict[str, Any],
annotated_function: Callable[[int, Literal["bar", "baz"]], None],
dummy_pydantic: type[BaseModel],
runnable: Runnable,
runnable: RunnableLambda[_Args, None],
dummy_typing_typed_dict: type,
dummy_typing_typed_dict_docstring: type,
dummy_extensions_typed_dict: type,
@@ -643,7 +644,7 @@ openai_function_no_description_no_params = {
openai_function_no_description,
],
)
def test_convert_to_openai_function_no_description(func: dict) -> None:
def test_convert_to_openai_function_no_description(func: dict[str, Any]) -> None:
expected = {
"name": "dummy_function",
"parameters": {
@@ -670,7 +671,9 @@ def test_convert_to_openai_function_no_description(func: dict) -> None:
openai_function_no_description_no_params,
],
)
def test_convert_to_openai_function_no_description_no_params(func: dict) -> None:
def test_convert_to_openai_function_no_description_no_params(
func: dict[str, Any],
) -> None:
expected = {
"name": "dummy_function",
}
@@ -1040,7 +1043,9 @@ def test__convert_typed_dict_to_openai_function(
@pytest.mark.parametrize("typed_dict", [ExtensionsTypedDict, TypingTypedDict])
def test__convert_typed_dict_to_openai_function_fail(typed_dict: type) -> None:
class Tool(typed_dict): # type: ignore[misc]
arg1: typing.MutableSet # Pydantic 2 supports this, but pydantic v1 does not.
arg1: typing.MutableSet[
Any
] # Pydantic 2 supports this, but pydantic v1 does not.
# Error should be raised since we're using v1 code path here
with pytest.raises(TypeError):
@@ -1081,15 +1086,15 @@ def test_convert_to_openai_function_no_args() -> None:
def test_convert_to_json_schema(
pydantic: type[BaseModel],
function: Callable,
function_docstring_annotations: Callable,
function: Callable[[int, Literal["bar", "baz"]], None],
function_docstring_annotations: Callable[[int, Literal["bar", "baz"]], None],
dummy_structured_tool: StructuredTool,
dummy_structured_tool_args_schema_dict: StructuredTool,
dummy_tool: BaseTool,
json_schema: dict,
anthropic_tool: dict,
bedrock_converse_tool: dict,
annotated_function: Callable,
json_schema: dict[str, Any],
anthropic_tool: dict[str, Any],
bedrock_converse_tool: dict[str, Any],
annotated_function: Callable[[int, Literal["bar", "baz"]], None],
dummy_pydantic: type[BaseModel],
dummy_typing_typed_dict: type,
dummy_typing_typed_dict_docstring: type,
@@ -1123,10 +1128,10 @@ def test_convert_to_json_schema(
def test_convert_to_openai_function_nested_strict_2() -> None:
def my_function(arg1: dict, arg2: dict | None) -> None:
def my_function(arg1: dict[str, Any], arg2: dict[str, Any] | None) -> None:
"""Dummy function."""
expected: dict = {
expected: dict[str, Any] = {
"name": "my_function",
"description": "Dummy function.",
"parameters": {
@@ -1,3 +1,5 @@
from typing import Any
import pytest
from langchain_core.utils.function_calling import _rm_titles
@@ -229,5 +231,5 @@ output5 = {
(schema5, output5),
],
)
def test_rm_titles(schema: dict, output: dict) -> None:
def test_rm_titles(schema: dict[str, Any], output: dict[str, Any]) -> None:
assert _rm_titles(schema) == output
@@ -132,7 +132,9 @@ def test_check_package_version(
],
)
def test_merge_dicts(
left: dict, right: dict, expected: dict | AbstractContextManager
left: dict[str, Any],
right: dict[str, Any],
expected: dict[str, Any] | AbstractContextManager[BaseException],
) -> None:
err = expected if isinstance(expected, AbstractContextManager) else nullcontext()
@@ -160,7 +162,9 @@ def test_merge_dicts(
)
@pytest.mark.xfail(reason="Refactors to make in 0.3")
def test_merge_dicts_0_3(
left: dict, right: dict, expected: dict | AbstractContextManager
left: dict[str, Any],
right: dict[str, Any],
expected: dict[str, Any] | AbstractContextManager[BaseException],
) -> None:
err = expected if isinstance(expected, AbstractContextManager) else nullcontext()
@@ -6,6 +6,7 @@ import pytest
pytest.importorskip("numpy")
import numpy as np
import numpy.typing as npt
from langchain_core.vectorstores.utils import _cosine_similarity
@@ -61,8 +62,8 @@ class TestCosineSimilarity:
def test_numpy_array_input(self) -> None:
"""Test with numpy array inputs."""
x: np.ndarray = np.array([[1, 0], [0, 1]])
y: np.ndarray = np.array([[1, 0], [0, 1]])
x: npt.NDArray[np.floating] = np.array([[1, 0], [0, 1]])
y: npt.NDArray[np.floating] = np.array([[1, 0], [0, 1]])
result = _cosine_similarity(x, y)
expected = np.array([[1.0, 0.0], [0.0, 1.0]])
np.testing.assert_array_almost_equal(result, expected)
@@ -70,7 +71,7 @@ class TestCosineSimilarity:
def test_mixed_input_types(self) -> None:
"""Test with mixed input types (list and numpy array)."""
x: list[list[float]] = [[1, 0], [0, 1]]
y: np.ndarray = np.array([[1, 0], [0, 1]])
y: npt.NDArray[np.floating] = np.array([[1, 0], [0, 1]])
result = _cosine_similarity(x, y)
expected = np.array([[1.0, 0.0], [0.0, 1.0]])
np.testing.assert_array_almost_equal(result, expected)
@@ -30,7 +30,7 @@ class CustomAddTextsVectorstore(VectorStore):
def add_texts(
self,
texts: Iterable[str],
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
ids: list[str] | None = None,
**kwargs: Any,
) -> list[str]:
@@ -58,7 +58,7 @@ class CustomAddTextsVectorstore(VectorStore):
cls,
texts: list[str],
embedding: Embeddings,
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
**kwargs: Any,
) -> CustomAddTextsVectorstore:
vectorstore = CustomAddTextsVectorstore()
@@ -104,7 +104,7 @@ class CustomAddDocumentsVectorstore(VectorStore):
cls,
texts: list[str],
embedding: Embeddings,
metadatas: list[dict] | None = None,
metadatas: list[dict[str, Any]] | None = None,
**kwargs: Any,
) -> CustomAddDocumentsVectorstore:
vectorstore = CustomAddDocumentsVectorstore()
@@ -455,7 +455,8 @@ class ExperimentalMarkdownSyntaxTextSplitter:
# Apply the header stack as metadata
for depth, value in self.current_header_stack:
header_key = self.splittable_headers.get("#" * depth)
self.current_chunk.metadata[header_key] = value
if header_key is not None:
self.current_chunk.metadata[header_key] = value
self.chunks.append(self.current_chunk)
# Reset the current chunk
self.current_chunk = Document(page_content="")