diff --git a/libs/core/langchain_core/_api/deprecation.py b/libs/core/langchain_core/_api/deprecation.py index 45930c3e6c..1492e21780 100644 --- a/libs/core/langchain_core/_api/deprecation.py +++ b/libs/core/langchain_core/_api/deprecation.py @@ -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" diff --git a/libs/core/langchain_core/agents.py b/libs/core/langchain_core/agents.py index 76f818b06a..1ccc1c1585 100644 --- a/libs/core/langchain_core/agents.py +++ b/libs/core/langchain_core/agents.py @@ -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) diff --git a/libs/core/langchain_core/callbacks/manager.py b/libs/core/langchain_core/callbacks/manager.py index c1ba9b76c0..40fd907272 100644 --- a/libs/core/langchain_core/callbacks/manager.py +++ b/libs/core/langchain_core/callbacks/manager.py @@ -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": diff --git a/libs/core/langchain_core/chat_sessions.py b/libs/core/langchain_core/chat_sessions.py index ed8c6343c5..c4219217ea 100644 --- a/libs/core/langchain_core/chat_sessions.py +++ b/libs/core/langchain_core/chat_sessions.py @@ -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.""" diff --git a/libs/core/langchain_core/document_loaders/langsmith.py b/libs/core/langchain_core/document_loaders/langsmith.py index 23a44e05d4..3c6a04b886 100644 --- a/libs/core/langchain_core/document_loaders/langsmith.py +++ b/libs/core/langchain_core/document_loaders/langsmith.py @@ -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, diff --git a/libs/core/langchain_core/documents/base.py b/libs/core/langchain_core/documents/base.py index 969ee49a17..efb2e23e90 100644 --- a/libs/core/langchain_core/documents/base.py +++ b/libs/core/langchain_core/documents/base.py @@ -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. diff --git a/libs/core/langchain_core/env.py b/libs/core/langchain_core/env.py index 240e62a8e6..20d384ee75 100644 --- a/libs/core/langchain_core/env.py +++ b/libs/core/langchain_core/env.py @@ -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: diff --git a/libs/core/langchain_core/example_selectors/base.py b/libs/core/langchain_core/example_selectors/base.py index ec845cfc2e..7297c75468 100644 --- a/libs/core/langchain_core/example_selectors/base.py +++ b/libs/core/langchain_core/example_selectors/base.py @@ -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: diff --git a/libs/core/langchain_core/example_selectors/length_based.py b/libs/core/langchain_core/example_selectors/length_based.py index e60e47e891..7205635f42 100644 --- a/libs/core/langchain_core/example_selectors/length_based.py +++ b/libs/core/langchain_core/example_selectors/length_based.py @@ -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. diff --git a/libs/core/langchain_core/example_selectors/semantic_similarity.py b/libs/core/langchain_core/example_selectors/semantic_similarity.py index 1e7491a2eb..42e711a8be 100644 --- a/libs/core/langchain_core/example_selectors/semantic_similarity.py +++ b/libs/core/langchain_core/example_selectors/semantic_similarity.py @@ -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. diff --git a/libs/core/langchain_core/language_models/_utils.py b/libs/core/langchain_core/language_models/_utils.py index 289b675307..c463db9c7b 100644 --- a/libs/core/langchain_core/language_models/_utils.py +++ b/libs/core/langchain_core/language_models/_utils.py @@ -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: diff --git a/libs/core/langchain_core/language_models/base.py b/libs/core/langchain_core/language_models/base.py index 570076290e..fb416887e7 100644 --- a/libs/core/langchain_core/language_models/base.py +++ b/libs/core/langchain_core/language_models/base.py @@ -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. diff --git a/libs/core/langchain_core/language_models/chat_models.py b/libs/core/langchain_core/language_models/chat_models.py index 76c8728c1b..cc946a919b 100644 --- a/libs/core/langchain_core/language_models/chat_models.py +++ b/libs/core/langchain_core/language_models/chat_models.py @@ -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, diff --git a/libs/core/langchain_core/language_models/llms.py b/libs/core/langchain_core/language_models/llms.py index 614e12ec16..79736b86e5 100644 --- a/libs/core/langchain_core/language_models/llms.py +++ b/libs/core/langchain_core/language_models/llms.py @@ -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 diff --git a/libs/core/langchain_core/load/serializable.py b/libs/core/langchain_core/load/serializable.py index 429a5e8f88..4764282655 100644 --- a/libs/core/langchain_core/load/serializable.py +++ b/libs/core/langchain_core/load/serializable.py @@ -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. diff --git a/libs/core/langchain_core/messages/ai.py b/libs/core/langchain_core/messages/ai.py index 92bac634d6..b6ed773872 100644 --- a/libs/core/langchain_core/messages/ai.py +++ b/libs/core/langchain_core/messages/ai.py @@ -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)), ), ) diff --git a/libs/core/langchain_core/messages/base.py b/libs/core/langchain_core/messages/base.py index 2b0e998c70..21760e2100 100644 --- a/libs/core/langchain_core/messages/base.py +++ b/libs/core/langchain_core/messages/base.py @@ -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: diff --git a/libs/core/langchain_core/messages/block_translators/anthropic.py b/libs/core/langchain_core/messages/block_translators/anthropic.py index eab2163f07..fb70388d60 100644 --- a/libs/core/langchain_core/messages/block_translators/anthropic.py +++ b/libs/core/langchain_core/messages/block_translators/anthropic.py @@ -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 diff --git a/libs/core/langchain_core/messages/block_translators/groq.py b/libs/core/langchain_core/messages/block_translators/groq.py index bcaa1a15b7..773f3fd7fb 100644 --- a/libs/core/langchain_core/messages/block_translators/groq.py +++ b/libs/core/langchain_core/messages/block_translators/groq.py @@ -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: diff --git a/libs/core/langchain_core/messages/block_translators/langchain_v0.py b/libs/core/langchain_core/messages/block_translators/langchain_v0.py index f7cb03839e..3d25e19f17 100644 --- a/libs/core/langchain_core/messages/block_translators/langchain_v0.py +++ b/libs/core/langchain_core/messages/block_translators/langchain_v0.py @@ -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: diff --git a/libs/core/langchain_core/messages/block_translators/openai.py b/libs/core/langchain_core/messages/block_translators/openai.py index 627459254c..86c201dcc4 100644 --- a/libs/core/langchain_core/messages/block_translators/openai.py +++ b/libs/core/langchain_core/messages/block_translators/openai.py @@ -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} diff --git a/libs/core/langchain_core/messages/content.py b/libs/core/langchain_core/messages/content.py index 3a02139d5b..58084e3f08 100644 --- a/libs/core/langchain_core/messages/content.py +++ b/libs/core/langchain_core/messages/content.py @@ -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. diff --git a/libs/core/langchain_core/messages/human.py b/libs/core/langchain_core/messages/human.py index 338e221370..c9a1d2756f 100644 --- a/libs/core/langchain_core/messages/human.py +++ b/libs/core/langchain_core/messages/human.py @@ -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: diff --git a/libs/core/langchain_core/messages/system.py b/libs/core/langchain_core/messages/system.py index 4a60811dff..3ace120f39 100644 --- a/libs/core/langchain_core/messages/system.py +++ b/libs/core/langchain_core/messages/system.py @@ -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: diff --git a/libs/core/langchain_core/messages/tool.py b/libs/core/langchain_core/messages/tool.py index a83d4e6eb9..7219942f02 100644 --- a/libs/core/langchain_core/messages/tool.py +++ b/libs/core/langchain_core/messages/tool.py @@ -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: diff --git a/libs/core/langchain_core/messages/utils.py b/libs/core/langchain_core/messages/utils.py index 095ae805b2..c1657145cc 100644 --- a/libs/core/langchain_core/messages/utils.py +++ b/libs/core/langchain_core/messages/utils.py @@ -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": [..., ""], - "kwargs": {...}}` — unpacked structurally and routed through the - standard dict-with-type dispatch. + `{"lc": 1, "type": "constructor", "id": [..., ""], + "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: diff --git a/libs/core/langchain_core/output_parsers/base.py b/libs/core/langchain_core/output_parsers/base.py index b316abad92..9f9548edbf 100644 --- a/libs/core/langchain_core/output_parsers/base.py +++ b/libs/core/langchain_core/output_parsers/base.py @@ -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. diff --git a/libs/core/langchain_core/output_parsers/list.py b/libs/core/langchain_core/output_parsers/list.py index 834c9ec153..f2c9cb5d45 100644 --- a/libs/core/langchain_core/output_parsers/list.py +++ b/libs/core/langchain_core/output_parsers/list.py @@ -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 diff --git a/libs/core/langchain_core/output_parsers/openai_tools.py b/libs/core/langchain_core/output_parsers/openai_tools.py index c42e094665..78be65ecef 100644 --- a/libs/core/langchain_core/output_parsers/openai_tools.py +++ b/libs/core/langchain_core/output_parsers/openai_tools.py @@ -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, diff --git a/libs/core/langchain_core/output_parsers/pydantic.py b/libs/core/langchain_core/output_parsers/pydantic.py index 7a7eee972d..ab89c91754 100644 --- a/libs/core/langchain_core/output_parsers/pydantic.py +++ b/libs/core/langchain_core/output_parsers/pydantic.py @@ -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__ diff --git a/libs/core/langchain_core/output_parsers/xml.py b/libs/core/langchain_core/output_parsers/xml.py index c65a1db329..a26a37ffcf 100644 --- a/libs/core/langchain_core/output_parsers/xml.py +++ b/libs/core/langchain_core/output_parsers/xml.py @@ -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): `'\n \n \n'` """ - 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: diff --git a/libs/core/langchain_core/outputs/chat_result.py b/libs/core/langchain_core/outputs/chat_result.py index 1cc814310e..2b7bc91fe5 100644 --- a/libs/core/langchain_core/outputs/chat_result.py +++ b/libs/core/langchain_core/outputs/chat_result.py @@ -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 diff --git a/libs/core/langchain_core/outputs/llm_result.py b/libs/core/langchain_core/outputs/llm_result.py index df40c41975..5e0fa42968 100644 --- a/libs/core/langchain_core/outputs/llm_result.py +++ b/libs/core/langchain_core/outputs/llm_result.py @@ -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 diff --git a/libs/core/langchain_core/prompt_values.py b/libs/core/langchain_core/prompt_values.py index e85fe1efb4..668871e221 100644 --- a/libs/core/langchain_core/prompt_values.py +++ b/libs/core/langchain_core/prompt_values.py @@ -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): diff --git a/libs/core/langchain_core/prompts/base.py b/libs/core/langchain_core/prompts/base.py index 9a9cc30242..e90220b796 100644 --- a/libs/core/langchain_core/prompts/base.py +++ b/libs/core/langchain_core/prompts/base.py @@ -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: diff --git a/libs/core/langchain_core/prompts/chat.py b/libs/core/langchain_core/prompts/chat.py index ebd58c8031..4a912f64b6 100644 --- a/libs/core/langchain_core/prompts/chat.py +++ b/libs/core/langchain_core/prompts/chat.py @@ -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. diff --git a/libs/core/langchain_core/prompts/dict.py b/libs/core/langchain_core/prompts/dict.py index 5a665bfbb8..bb68d4f28d 100644 --- a/libs/core/langchain_core/prompts/dict.py +++ b/libs/core/langchain_core/prompts/dict.py @@ -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(): diff --git a/libs/core/langchain_core/prompts/few_shot.py b/libs/core/langchain_core/prompts/few_shot.py index 8e8e9aa315..cfe46d84d2 100644 --- a/libs/core/langchain_core/prompts/few_shot.py +++ b/libs/core/langchain_core/prompts/few_shot.py @@ -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: diff --git a/libs/core/langchain_core/prompts/few_shot_with_templates.py b/libs/core/langchain_core/prompts/few_shot_with_templates.py index ca664cabee..f70cdc6e66 100644 --- a/libs/core/langchain_core/prompts/few_shot_with_templates.py +++ b/libs/core/langchain_core/prompts/few_shot_with_templates.py @@ -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: diff --git a/libs/core/langchain_core/prompts/image.py b/libs/core/langchain_core/prompts/image.py index ee8c9421f2..ce8054bf59 100644 --- a/libs/core/langchain_core/prompts/image.py +++ b/libs/core/langchain_core/prompts/image.py @@ -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" diff --git a/libs/core/langchain_core/prompts/loading.py b/libs/core/langchain_core/prompts/loading.py index d130f9d871..bfd2f23856 100644 --- a/libs/core/langchain_core/prompts/loading.py +++ b/libs/core/langchain_core/prompts/loading.py @@ -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, diff --git a/libs/core/langchain_core/prompts/prompt.py b/libs/core/langchain_core/prompts/prompt.py index cef55a5c2f..bbf72ebc00 100644 --- a/libs/core/langchain_core/prompts/prompt.py +++ b/libs/core/langchain_core/prompts/prompt.py @@ -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 diff --git a/libs/core/langchain_core/prompts/string.py b/libs/core/langchain_core/prompts/string.py index f37bdba221..f96d83e492 100644 --- a/libs/core/langchain_core/prompts/string.py +++ b/libs/core/langchain_core/prompts/string.py @@ -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 diff --git a/libs/core/langchain_core/prompts/structured.py b/libs/core/langchain_core/prompts/structured.py index 00ac407fb7..0150007984 100644 --- a/libs/core/langchain_core/prompts/structured.py +++ b/libs/core/langchain_core/prompts/structured.py @@ -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: diff --git a/libs/core/langchain_core/runnables/base.py b/libs/core/langchain_core/runnables/base.py index 3a9b0cdcfc..65a53fb1d6 100644 --- a/libs/core/langchain_core/runnables/base.py +++ b/libs/core/langchain_core/runnables/base.py @@ -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: diff --git a/libs/core/langchain_core/runnables/branch.py b/libs/core/langchain_core/runnables/branch.py index ca3dfd99da..b8a3ffb7d1 100644 --- a/libs/core/langchain_core/runnables/branch.py +++ b/libs/core/langchain_core/runnables/branch.py @@ -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_ = [] diff --git a/libs/core/langchain_core/runnables/graph.py b/libs/core/langchain_core/runnables/graph.py index cdab7d4884..4b5bfdfb83 100644 --- a/libs/core/langchain_core/runnables/graph.py +++ b/libs/core/langchain_core/runnables/graph.py @@ -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, diff --git a/libs/core/langchain_core/runnables/history.py b/libs/core/langchain_core/runnables/history.py index c85386735c..d04e3226d9 100644 --- a/libs/core/langchain_core/runnables/history.py +++ b/libs/core/langchain_core/runnables/history.py @@ -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): diff --git a/libs/core/langchain_core/runnables/passthrough.py b/libs/core/langchain_core/runnables/passthrough.py index f5e01cfe20..0df584d252 100644 --- a/libs/core/langchain_core/runnables/passthrough.py +++ b/libs/core/langchain_core/runnables/passthrough.py @@ -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( diff --git a/libs/core/langchain_core/runnables/utils.py b/libs/core/langchain_core/runnables/utils.py index e46251a207..1cf2371e63 100644 --- a/libs/core/langchain_core/runnables/utils.py +++ b/libs/core/langchain_core/runnables/utils.py @@ -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: diff --git a/libs/core/langchain_core/tools/base.py b/libs/core/langchain_core/tools/base.py index 67c64fb1f1..a04f0abef4 100644 --- a/libs/core/langchain_core/tools/base.py +++ b/libs/core/langchain_core/tools/base.py @@ -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( diff --git a/libs/core/langchain_core/tools/convert.py b/libs/core/langchain_core/tools/convert.py index 48c518a298..5781afc473 100644 --- a/libs/core/langchain_core/tools/convert.py +++ b/libs/core/langchain_core/tools/convert.py @@ -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, diff --git a/libs/core/langchain_core/tools/retriever.py b/libs/core/langchain_core/tools/retriever.py index 9e2d84dcb0..890e4ca0ff 100644 --- a/libs/core/langchain_core/tools/retriever.py +++ b/libs/core/langchain_core/tools/retriever.py @@ -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: diff --git a/libs/core/langchain_core/tools/simple.py b/libs/core/langchain_core/tools/simple.py index ca80164df8..502960ce1f 100644 --- a/libs/core/langchain_core/tools/simple.py +++ b/libs/core/langchain_core/tools/simple.py @@ -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 diff --git a/libs/core/langchain_core/tools/structured.py b/libs/core/langchain_core/tools/structured.py index 8b67e3b454..e9643a0ad8 100644 --- a/libs/core/langchain_core/tools/structured.py +++ b/libs/core/langchain_core/tools/structured.py @@ -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) diff --git a/libs/core/langchain_core/tracers/base.py b/libs/core/langchain_core/tracers/base.py index b52420f0d8..0c49530a30 100644 --- a/libs/core/langchain_core/tracers/base.py +++ b/libs/core/langchain_core/tracers/base.py @@ -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 diff --git a/libs/core/langchain_core/tracers/core.py b/libs/core/langchain_core/tracers/core.py index 75614e3c88..c6f6cca167 100644 --- a/libs/core/langchain_core/tracers/core.py +++ b/libs/core/langchain_core/tracers/core.py @@ -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 diff --git a/libs/core/langchain_core/tracers/evaluation.py b/libs/core/langchain_core/tracers/evaluation.py index 22c6f600f5..a7062cb2a2 100644 --- a/libs/core/langchain_core/tracers/evaluation.py +++ b/libs/core/langchain_core/tracers/evaluation.py @@ -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 diff --git a/libs/core/langchain_core/tracers/event_stream.py b/libs/core/langchain_core/tracers/event_stream.py index 399a7c19b6..cba55b8bb9 100644 --- a/libs/core/langchain_core/tracers/event_stream.py +++ b/libs/core/langchain_core/tracers/event_stream.py @@ -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 diff --git a/libs/core/langchain_core/tracers/log_stream.py b/libs/core/langchain_core/tracers/log_stream.py index 5131815ebd..ebafc14b0e 100644 --- a/libs/core/langchain_core/tracers/log_stream.py +++ b/libs/core/langchain_core/tracers/log_stream.py @@ -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__( diff --git a/libs/core/langchain_core/tracers/memory_stream.py b/libs/core/langchain_core/tracers/memory_stream.py index 42e74fb00d..9047e60105 100644 --- a/libs/core/langchain_core/tracers/memory_stream.py +++ b/libs/core/langchain_core/tracers/memory_stream.py @@ -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]: diff --git a/libs/core/langchain_core/utils/_merge.py b/libs/core/langchain_core/utils/_merge.py index 6a0cb38f07..d6e0cce4db 100644 --- a/libs/core/langchain_core/utils/_merge.py +++ b/libs/core/langchain_core/utils/_merge.py @@ -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: diff --git a/libs/core/langchain_core/utils/aiter.py b/libs/core/langchain_core/utils/aiter.py index e5dc0d1aea..912fb32e70 100644 --- a/libs/core/langchain_core/utils/aiter.py +++ b/libs/core/langchain_core/utils/aiter.py @@ -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: diff --git a/libs/core/langchain_core/utils/formatting.py b/libs/core/langchain_core/utils/formatting.py index 48905a4cc0..2458fdcf4c 100644 --- a/libs/core/langchain_core/utils/formatting.py +++ b/libs/core/langchain_core/utils/formatting.py @@ -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. diff --git a/libs/core/langchain_core/utils/function_calling.py b/libs/core/langchain_core/utils/function_calling.py index c66df94437..eae4a8c8a0 100644 --- a/libs/core/langchain_core/utils/function_calling.py +++ b/libs/core/langchain_core/utils/function_calling.py @@ -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. diff --git a/libs/core/langchain_core/utils/html.py b/libs/core/langchain_core/utils/html.py index 4798b02ce7..a2f9424e38 100644 --- a/libs/core/langchain_core/utils/html.py +++ b/libs/core/langchain_core/utils/html.py @@ -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, diff --git a/libs/core/langchain_core/utils/input.py b/libs/core/langchain_core/utils/input.py index d97d4006d3..34eb2c2bb0 100644 --- a/libs/core/langchain_core/utils/input.py +++ b/libs/core/langchain_core/utils/input.py @@ -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. diff --git a/libs/core/langchain_core/utils/json.py b/libs/core/langchain_core/utils/json.py index a836ffc4e6..d772a81078 100644 --- a/libs/core/langchain_core/utils/json.py +++ b/libs/core/langchain_core/utils/json.py @@ -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. diff --git a/libs/core/langchain_core/utils/json_schema.py b/libs/core/langchain_core/utils/json_schema.py index d1ff1de5fc..eface5da2e 100644 --- a/libs/core/langchain_core/utils/json_schema.py +++ b/libs/core/langchain_core/utils/json_schema.py @@ -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 ), diff --git a/libs/core/langchain_core/utils/pydantic.py b/libs/core/langchain_core/utils/pydantic.py index d1c152d8d5..7fed428145 100644 --- a/libs/core/langchain_core/utils/pydantic.py +++ b/libs/core/langchain_core/utils/pydantic.py @@ -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. diff --git a/libs/core/langchain_core/utils/strings.py b/libs/core/langchain_core/utils/strings.py index 357b16f8e1..c9f001934f 100644 --- a/libs/core/langchain_core/utils/strings.py +++ b/libs/core/langchain_core/utils/strings.py @@ -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: diff --git a/libs/core/langchain_core/utils/usage.py b/libs/core/langchain_core/utils/usage.py index 47e483a555..99784f2c6c 100644 --- a/libs/core/langchain_core/utils/usage.py +++ b/libs/core/langchain_core/utils/usage.py @@ -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 diff --git a/libs/core/langchain_core/utils/utils.py b/libs/core/langchain_core/utils/utils.py index e8a5ed999a..ac9c3e0611 100644 --- a/libs/core/langchain_core/utils/utils.py +++ b/libs/core/langchain_core/utils/utils.py @@ -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.""" diff --git a/libs/core/langchain_core/vectorstores/base.py b/libs/core/langchain_core/vectorstores/base.py index 827a05cc90..7d14506182 100644 --- a/libs/core/langchain_core/vectorstores/base.py +++ b/libs/core/langchain_core/vectorstores/base.py @@ -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: diff --git a/libs/core/langchain_core/vectorstores/in_memory.py b/libs/core/langchain_core/vectorstores/in_memory.py index ef3c78ab60..651db514d5 100644 --- a/libs/core/langchain_core/vectorstores/in_memory.py +++ b/libs/core/langchain_core/vectorstores/in_memory.py @@ -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( diff --git a/libs/core/langchain_core/vectorstores/utils.py b/libs/core/langchain_core/vectorstores/utils.py index 551524beb3..04914e86af 100644 --- a/libs/core/langchain_core/vectorstores/utils.py +++ b/libs/core/langchain_core/vectorstores/utils.py @@ -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]: diff --git a/libs/core/tests/unit_tests/documents/test_document.py b/libs/core/tests/unit_tests/documents/test_document.py index e312121bd0..9392d746ba 100644 --- a/libs/core/tests/unit_tests/documents/test_document.py +++ b/libs/core/tests/unit_tests/documents/test_document.py @@ -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 diff --git a/libs/core/tests/unit_tests/example_selectors/test_similarity.py b/libs/core/tests/unit_tests/example_selectors/test_similarity.py index b07870121c..408c12e653 100644 --- a/libs/core/tests/unit_tests/example_selectors/test_similarity.py +++ b/libs/core/tests/unit_tests/example_selectors/test_similarity.py @@ -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) diff --git a/libs/core/tests/unit_tests/messages/test_ai.py b/libs/core/tests/unit_tests/messages/test_ai.py index 31f8b3e1bb..a9c732c99d 100644 --- a/libs/core/tests/unit_tests/messages/test_ai.py +++ b/libs/core/tests/unit_tests/messages/test_ai.py @@ -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 == [ { diff --git a/libs/core/tests/unit_tests/messages/test_utils.py b/libs/core/tests/unit_tests/messages/test_utils.py index 9b4123d859..07846eb1c2 100644 --- a/libs/core/tests/unit_tests/messages/test_utils.py +++ b/libs/core/tests/unit_tests/messages/test_utils.py @@ -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", diff --git a/libs/core/tests/unit_tests/prompts/test_structured.py b/libs/core/tests/unit_tests/prompts/test_structured.py index 77157dfae1..5cd109dcce 100644 --- a/libs/core/tests/unit_tests/prompts/test_structured.py +++ b/libs/core/tests/unit_tests/prompts/test_structured.py @@ -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 diff --git a/libs/core/tests/unit_tests/pydantic_utils.py b/libs/core/tests/unit_tests/pydantic_utils.py index 2f01494e1d..1379c53d52 100644 --- a/libs/core/tests/unit_tests/pydantic_utils.py +++ b/libs/core/tests/unit_tests/pydantic_utils.py @@ -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 diff --git a/libs/core/tests/unit_tests/runnables/test_fallbacks.py b/libs/core/tests/unit_tests/runnables/test_fallbacks.py index 01b6b1a363..11e02f3939 100644 --- a/libs/core/tests/unit_tests/runnables/test_fallbacks.py +++ b/libs/core/tests/unit_tests/runnables/test_fallbacks.py @@ -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) diff --git a/libs/core/tests/unit_tests/runnables/test_graph.py b/libs/core/tests/unit_tests/runnables/test_graph.py index c7d705cfd6..ffd90a0779 100644 --- a/libs/core/tests/unit_tests/runnables/test_graph.py +++ b/libs/core/tests/unit_tests/runnables/test_graph.py @@ -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 diff --git a/libs/core/tests/unit_tests/runnables/test_runnable.py b/libs/core/tests/unit_tests/runnables/test_runnable.py index da648b0031..1cccc2743b 100644 --- a/libs/core/tests/unit_tests/runnables/test_runnable.py +++ b/libs/core/tests/unit_tests/runnables/test_runnable.py @@ -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) diff --git a/libs/core/tests/unit_tests/runnables/test_runnable_events_v1.py b/libs/core/tests/unit_tests/runnables/test_runnable_events_v1.py index 06f13e035b..1165382399 100644 --- a/libs/core/tests/unit_tests/runnables/test_runnable_events_v1.py +++ b/libs/core/tests/unit_tests/runnables/test_runnable_events_v1.py @@ -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} diff --git a/libs/core/tests/unit_tests/runnables/test_runnable_events_v2.py b/libs/core/tests/unit_tests/runnables/test_runnable_events_v2.py index f36dc87cbd..f59500dba4 100644 --- a/libs/core/tests/unit_tests/runnables/test_runnable_events_v2.py +++ b/libs/core/tests/unit_tests/runnables/test_runnable_events_v2.py @@ -93,7 +93,7 @@ async def _collect_events( async def test_event_stream_with_simple_function_tool() -> None: """Test the event stream with a function and tool.""" - def foo(x: int) -> dict: + 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} diff --git a/libs/core/tests/unit_tests/runnables/test_tracing_interops.py b/libs/core/tests/unit_tests/runnables/test_tracing_interops.py index 029e3bcb8e..c1f11d9c2a 100644 --- a/libs/core/tests/unit_tests/runnables/test_tracing_interops.py +++ b/libs/core/tests/unit_tests/runnables/test_tracing_interops.py @@ -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]: diff --git a/libs/core/tests/unit_tests/runnables/test_utils.py b/libs/core/tests/unit_tests/runnables/test_utils.py index 37c19ca1a1..031bb5aa23 100644 --- a/libs/core/tests/unit_tests/runnables/test_utils.py +++ b/libs/core/tests/unit_tests/runnables/test_utils.py @@ -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: diff --git a/libs/core/tests/unit_tests/stores/test_in_memory.py b/libs/core/tests/unit_tests/stores/test_in_memory.py index 6c24ebe463..4602777b85 100644 --- a/libs/core/tests/unit_tests/stores/test_in_memory.py +++ b/libs/core/tests/unit_tests/stores/test_in_memory.py @@ -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: diff --git a/libs/core/tests/unit_tests/test_messages.py b/libs/core/tests/unit_tests/test_messages.py index a13fdc62a6..c6e4f5b50b 100644 --- a/libs/core/tests/unit_tests/test_messages.py +++ b/libs/core/tests/unit_tests/test_messages.py @@ -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 diff --git a/libs/core/tests/unit_tests/test_prompt_values.py b/libs/core/tests/unit_tests/test_prompt_values.py index 6a08a4270a..187ca9f7b3 100644 --- a/libs/core/tests/unit_tests/test_prompt_values.py +++ b/libs/core/tests/unit_tests/test_prompt_values.py @@ -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"), diff --git a/libs/core/tests/unit_tests/test_tools.py b/libs/core/tests/unit_tests/test_tools.py index 274b25874f..43c23a109b 100644 --- a/libs/core/tests/unit_tests/test_tools.py +++ b/libs/core/tests/unit_tests/test_tools.py @@ -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. diff --git a/libs/core/tests/unit_tests/tracers/test_memory_stream.py b/libs/core/tests/unit_tests/tracers/test_memory_stream.py index ba9daa913c..c0840a3c19 100644 --- a/libs/core/tests/unit_tests/tracers/test_memory_stream.py +++ b/libs/core/tests/unit_tests/tracers/test_memory_stream.py @@ -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() diff --git a/libs/core/tests/unit_tests/utils/test_function_calling.py b/libs/core/tests/unit_tests/utils/test_function_calling.py index d2906f59e3..e6e46fd450 100644 --- a/libs/core/tests/unit_tests/utils/test_function_calling.py +++ b/libs/core/tests/unit_tests/utils/test_function_calling.py @@ -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": { diff --git a/libs/core/tests/unit_tests/utils/test_rm_titles.py b/libs/core/tests/unit_tests/utils/test_rm_titles.py index ead63cfab4..285ce35987 100644 --- a/libs/core/tests/unit_tests/utils/test_rm_titles.py +++ b/libs/core/tests/unit_tests/utils/test_rm_titles.py @@ -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 diff --git a/libs/core/tests/unit_tests/utils/test_utils.py b/libs/core/tests/unit_tests/utils/test_utils.py index 815296f888..7799c9c9df 100644 --- a/libs/core/tests/unit_tests/utils/test_utils.py +++ b/libs/core/tests/unit_tests/utils/test_utils.py @@ -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() diff --git a/libs/core/tests/unit_tests/vectorstores/test_utils.py b/libs/core/tests/unit_tests/vectorstores/test_utils.py index 2ff9817f5b..d97348c347 100644 --- a/libs/core/tests/unit_tests/vectorstores/test_utils.py +++ b/libs/core/tests/unit_tests/vectorstores/test_utils.py @@ -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) diff --git a/libs/core/tests/unit_tests/vectorstores/test_vectorstore.py b/libs/core/tests/unit_tests/vectorstores/test_vectorstore.py index f8af00ee79..5ba467767f 100644 --- a/libs/core/tests/unit_tests/vectorstores/test_vectorstore.py +++ b/libs/core/tests/unit_tests/vectorstores/test_vectorstore.py @@ -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() diff --git a/libs/text-splitters/langchain_text_splitters/markdown.py b/libs/text-splitters/langchain_text_splitters/markdown.py index f4c2b05530..7eee94e9c5 100644 --- a/libs/text-splitters/langchain_text_splitters/markdown.py +++ b/libs/text-splitters/langchain_text_splitters/markdown.py @@ -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="")