langchain-litellm 0.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2024 LangChain, Inc.
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -0,0 +1,56 @@
1
+ Metadata-Version: 2.3
2
+ Name: langchain-litellm
3
+ Version: 0.1.0
4
+ Summary: An integration package connecting Litellm and LangChain
5
+ License: MIT
6
+ Requires-Python: >=3.9,<4.0
7
+ Classifier: License :: OSI Approved :: MIT License
8
+ Classifier: Programming Language :: Python :: 3
9
+ Classifier: Programming Language :: Python :: 3.9
10
+ Classifier: Programming Language :: Python :: 3.10
11
+ Classifier: Programming Language :: Python :: 3.11
12
+ Classifier: Programming Language :: Python :: 3.12
13
+ Classifier: Programming Language :: Python :: 3.13
14
+ Requires-Dist: langchain-core (>=0.3.15,<0.4.0)
15
+ Requires-Dist: litellm (>=1.65.1,<2.0.0)
16
+ Project-URL: Repository, https://github.com/langchain-ai/langchain
17
+ Project-URL: Release Notes, https://github.com/langchain-ai/langchain/releases?q=tag%3A%22litellm%3D%3D0%22&expanded=true
18
+ Project-URL: Source Code, https://github.com/langchain-ai/langchain/tree/master/libs/partners/litellm
19
+ Description-Content-Type: text/markdown
20
+
21
+ # langchain-litellm
22
+
23
+ This package contains the [LangChain](https://github.com/langchain-ai/langchain) integration with [LiteLLM](https://github.com/BerriAI/litellm)
24
+
25
+ ## Installation
26
+
27
+ ```bash
28
+ pip install -qU langchain-litellm
29
+ ```
30
+
31
+ ## Chat Models
32
+
33
+ `ChatLiteLLM` class exposes chat models from [LiteLLM](https://github.com/BerriAI/litellm).
34
+
35
+ ```python
36
+ from langchain_litellm.chat_models import ChatLiteLLM
37
+ from langchain_core.messages import HumanMessage
38
+ messages = [
39
+ HumanMessage(
40
+ content="Translate this sentence from English to French. I love programming."
41
+ )
42
+ ]
43
+ chat(messages)
44
+ ```
45
+
46
+ ## `ChatLiteLLM` also supports async and streaming functionality:
47
+ ```python
48
+ from langchain_core.callbacks import CallbackManager, StreamingStdOutCallbackHandler
49
+ await chat.agenerate([messages])
50
+ chat = ChatLiteLLM(
51
+ streaming=True,
52
+ verbose=True,
53
+ callback_manager=CallbackManager([StreamingStdOutCallbackHandler()]),
54
+ )
55
+ chat(messages)
56
+ ```
@@ -0,0 +1,36 @@
1
+ # langchain-litellm
2
+
3
+ This package contains the [LangChain](https://github.com/langchain-ai/langchain) integration with [LiteLLM](https://github.com/BerriAI/litellm)
4
+
5
+ ## Installation
6
+
7
+ ```bash
8
+ pip install -qU langchain-litellm
9
+ ```
10
+
11
+ ## Chat Models
12
+
13
+ `ChatLiteLLM` class exposes chat models from [LiteLLM](https://github.com/BerriAI/litellm).
14
+
15
+ ```python
16
+ from langchain_litellm.chat_models import ChatLiteLLM
17
+ from langchain_core.messages import HumanMessage
18
+ messages = [
19
+ HumanMessage(
20
+ content="Translate this sentence from English to French. I love programming."
21
+ )
22
+ ]
23
+ chat(messages)
24
+ ```
25
+
26
+ ## `ChatLiteLLM` also supports async and streaming functionality:
27
+ ```python
28
+ from langchain_core.callbacks import CallbackManager, StreamingStdOutCallbackHandler
29
+ await chat.agenerate([messages])
30
+ chat = ChatLiteLLM(
31
+ streaming=True,
32
+ verbose=True,
33
+ callback_manager=CallbackManager([StreamingStdOutCallbackHandler()]),
34
+ )
35
+ chat(messages)
36
+ ```
@@ -0,0 +1,16 @@
1
+ from importlib import metadata
2
+
3
+ from langchain_litellm.chat_models import ChatLiteLLM
4
+
5
+
6
+ try:
7
+ __version__ = metadata.version(__package__)
8
+ except metadata.PackageNotFoundError:
9
+ # Case where package metadata is not available.
10
+ __version__ = ""
11
+ del metadata # optional, avoids polluting the results of dir(__package__)
12
+
13
+ __all__ = [
14
+ "ChatLiteLLM",
15
+ "__version__",
16
+ ]
@@ -0,0 +1,603 @@
1
+ """Wrapper around LiteLLM's model I/O library."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import logging
7
+ from typing import (
8
+ Any,
9
+ AsyncIterator,
10
+ Callable,
11
+ Dict,
12
+ Iterator,
13
+ List,
14
+ Literal,
15
+ Mapping,
16
+ Optional,
17
+ Sequence,
18
+ Tuple,
19
+ Type,
20
+ Union,
21
+ )
22
+
23
+ from langchain_core.callbacks import (
24
+ AsyncCallbackManagerForLLMRun,
25
+ CallbackManagerForLLMRun,
26
+ )
27
+ from langchain_core.language_models import LanguageModelInput
28
+ from langchain_core.language_models.chat_models import (
29
+ BaseChatModel,
30
+ agenerate_from_stream,
31
+ generate_from_stream,
32
+ )
33
+ from langchain_core.language_models.llms import create_base_retry_decorator
34
+ from langchain_core.messages import (
35
+ AIMessage,
36
+ AIMessageChunk,
37
+ BaseMessage,
38
+ BaseMessageChunk,
39
+ ChatMessage,
40
+ ChatMessageChunk,
41
+ FunctionMessage,
42
+ FunctionMessageChunk,
43
+ HumanMessage,
44
+ HumanMessageChunk,
45
+ SystemMessage,
46
+ SystemMessageChunk,
47
+ ToolCall,
48
+ ToolCallChunk,
49
+ ToolMessage,
50
+ )
51
+ from langchain_core.messages.ai import UsageMetadata
52
+ from langchain_core.outputs import (
53
+ ChatGeneration,
54
+ ChatGenerationChunk,
55
+ ChatResult,
56
+ )
57
+ from langchain_core.runnables import Runnable
58
+ from langchain_core.tools import BaseTool
59
+ from langchain_core.utils import get_from_dict_or_env, pre_init
60
+ from langchain_core.utils.function_calling import convert_to_openai_tool
61
+ from pydantic import BaseModel, Field
62
+
63
+ logger = logging.getLogger(__name__)
64
+
65
+
66
+ class ChatLiteLLMException(Exception):
67
+ """Error with the `LiteLLM I/O` library"""
68
+
69
+
70
+ def _create_retry_decorator(
71
+ llm: ChatLiteLLM,
72
+ run_manager: Optional[
73
+ Union[AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun]
74
+ ] = None,
75
+ ) -> Callable[[Any], Any]:
76
+ """Returns a tenacity retry decorator, preconfigured to handle PaLM exceptions"""
77
+ import litellm
78
+
79
+ errors = [
80
+ litellm.Timeout,
81
+ litellm.APIError,
82
+ litellm.APIConnectionError,
83
+ litellm.RateLimitError,
84
+ ]
85
+ return create_base_retry_decorator(
86
+ error_types=errors, max_retries=llm.max_retries, run_manager=run_manager
87
+ )
88
+
89
+
90
+ def _convert_dict_to_message(_dict: Mapping[str, Any]) -> BaseMessage:
91
+ role = _dict["role"]
92
+ if role == "user":
93
+ return HumanMessage(content=_dict["content"])
94
+ elif role == "assistant":
95
+ # Fix for azure
96
+ # Also OpenAI returns None for tool invocations
97
+ content = _dict.get("content", "") or ""
98
+
99
+ additional_kwargs = {}
100
+ if _dict.get("function_call"):
101
+ additional_kwargs["function_call"] = dict(_dict["function_call"])
102
+
103
+ if _dict.get("tool_calls"):
104
+ additional_kwargs["tool_calls"] = _dict["tool_calls"]
105
+
106
+ return AIMessage(content=content, additional_kwargs=additional_kwargs)
107
+ elif role == "system":
108
+ return SystemMessage(content=_dict["content"])
109
+ elif role == "function":
110
+ return FunctionMessage(content=_dict["content"], name=_dict["name"])
111
+ else:
112
+ return ChatMessage(content=_dict["content"], role=role)
113
+
114
+
115
+ async def acompletion_with_retry(
116
+ llm: ChatLiteLLM,
117
+ run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
118
+ **kwargs: Any,
119
+ ) -> Any:
120
+ """Use tenacity to retry the async completion call."""
121
+ retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)
122
+
123
+ @retry_decorator
124
+ async def _completion_with_retry(**kwargs: Any) -> Any:
125
+ # Use OpenAI's async api https://github.com/openai/openai-python#async-api
126
+ return await llm.client.acreate(**kwargs)
127
+
128
+ return await _completion_with_retry(**kwargs)
129
+
130
+
131
+ def _convert_delta_to_message_chunk(
132
+ _dict: Mapping[str, Any], default_class: Type[BaseMessageChunk]
133
+ ) -> BaseMessageChunk:
134
+ role = _dict.get("role")
135
+ content = _dict.get("content") or ""
136
+ if _dict.get("function_call"):
137
+ additional_kwargs = {"function_call": dict(_dict["function_call"])}
138
+ elif _dict.get("reasoning_content"):
139
+ additional_kwargs = {"reasoning_content": _dict["reasoning_content"]}
140
+ else:
141
+ additional_kwargs = {}
142
+
143
+ tool_call_chunks = []
144
+ if raw_tool_calls := _dict.get("tool_calls"):
145
+ additional_kwargs["tool_calls"] = raw_tool_calls
146
+ try:
147
+ tool_call_chunks = [
148
+ ToolCallChunk(
149
+ name=rtc["function"].get("name"),
150
+ args=rtc["function"].get("arguments"),
151
+ id=rtc.get("id"),
152
+ index=rtc["index"],
153
+ )
154
+ for rtc in raw_tool_calls
155
+ ]
156
+ except KeyError:
157
+ pass
158
+
159
+ if role == "user" or default_class == HumanMessageChunk:
160
+ return HumanMessageChunk(content=content)
161
+ elif role == "assistant" or default_class == AIMessageChunk:
162
+ return AIMessageChunk(
163
+ content=content,
164
+ additional_kwargs=additional_kwargs,
165
+ tool_call_chunks=tool_call_chunks,
166
+ )
167
+ elif role == "system" or default_class == SystemMessageChunk:
168
+ return SystemMessageChunk(content=content)
169
+ elif role == "function" or default_class == FunctionMessageChunk:
170
+ return FunctionMessageChunk(content=content, name=_dict["name"])
171
+ elif role or default_class == ChatMessageChunk:
172
+ return ChatMessageChunk(content=content, role=role) # type: ignore[arg-type]
173
+ else:
174
+ return default_class(content=content) # type: ignore[call-arg]
175
+
176
+
177
+ def _lc_tool_call_to_openai_tool_call(tool_call: ToolCall) -> dict:
178
+ return {
179
+ "type": "function",
180
+ "id": tool_call["id"],
181
+ "function": {
182
+ "name": tool_call["name"],
183
+ "arguments": json.dumps(tool_call["args"]),
184
+ },
185
+ }
186
+
187
+
188
+ def _convert_message_to_dict(message: BaseMessage) -> dict:
189
+ message_dict: Dict[str, Any] = {"content": message.content}
190
+ if isinstance(message, ChatMessage):
191
+ message_dict["role"] = message.role
192
+ elif isinstance(message, HumanMessage):
193
+ message_dict["role"] = "user"
194
+ elif isinstance(message, AIMessage):
195
+ message_dict["role"] = "assistant"
196
+ if "function_call" in message.additional_kwargs:
197
+ message_dict["function_call"] = message.additional_kwargs["function_call"]
198
+ if message.tool_calls:
199
+ message_dict["tool_calls"] = [
200
+ _lc_tool_call_to_openai_tool_call(tc) for tc in message.tool_calls
201
+ ]
202
+ elif "tool_calls" in message.additional_kwargs:
203
+ message_dict["tool_calls"] = message.additional_kwargs["tool_calls"]
204
+ elif isinstance(message, SystemMessage):
205
+ message_dict["role"] = "system"
206
+ elif isinstance(message, FunctionMessage):
207
+ message_dict["role"] = "function"
208
+ message_dict["name"] = message.name
209
+ elif isinstance(message, ToolMessage):
210
+ message_dict["role"] = "tool"
211
+ message_dict["tool_call_id"] = message.tool_call_id
212
+ else:
213
+ raise ValueError(f"Got unknown type {message}")
214
+ if "name" in message.additional_kwargs:
215
+ message_dict["name"] = message.additional_kwargs["name"]
216
+ return message_dict
217
+
218
+
219
+ _OPENAI_MODELS = [
220
+ "o1-mini",
221
+ "o1-preview",
222
+ "gpt-4o-mini",
223
+ "gpt-4o-mini-2024-07-18",
224
+ "gpt-4o",
225
+ "gpt-4o-2024-08-06",
226
+ "gpt-4o-2024-05-13",
227
+ "gpt-4-turbo",
228
+ "gpt-4-turbo-preview",
229
+ "gpt-4-0125-preview",
230
+ "gpt-4-1106-preview",
231
+ "gpt-3.5-turbo-1106",
232
+ "gpt-3.5-turbo",
233
+ "gpt-3.5-turbo-0301",
234
+ "gpt-3.5-turbo-0613",
235
+ "gpt-3.5-turbo-16k",
236
+ "gpt-3.5-turbo-16k-0613",
237
+ "gpt-4",
238
+ "gpt-4-0314",
239
+ "gpt-4-0613",
240
+ "gpt-4-32k",
241
+ "gpt-4-32k-0314",
242
+ "gpt-4-32k-0613",
243
+ ]
244
+
245
+
246
+ class ChatLiteLLM(BaseChatModel):
247
+ """Chat model that uses the LiteLLM API."""
248
+
249
+ client: Any = None #: :meta private:
250
+ model: str = "gpt-3.5-turbo"
251
+ model_name: Optional[str] = None
252
+ """Model name to use."""
253
+ openai_api_key: Optional[str] = None
254
+ azure_api_key: Optional[str] = None
255
+ anthropic_api_key: Optional[str] = None
256
+ replicate_api_key: Optional[str] = None
257
+ cohere_api_key: Optional[str] = None
258
+ openrouter_api_key: Optional[str] = None
259
+ api_key: Optional[str] = None
260
+ streaming: bool = False
261
+ api_base: Optional[str] = None
262
+ organization: Optional[str] = None
263
+ custom_llm_provider: Optional[str] = None
264
+ request_timeout: Optional[Union[float, Tuple[float, float]]] = None
265
+ temperature: Optional[float] = None
266
+ """Run inference with this temperature. Must be in the closed
267
+ interval [0.0, 1.0]."""
268
+ model_kwargs: Dict[str, Any] = Field(default_factory=dict)
269
+ """Holds any model parameters valid for API call not explicitly specified."""
270
+ top_p: Optional[float] = None
271
+ """Decode using nucleus sampling: consider the smallest set of tokens whose
272
+ probability sum is at least top_p. Must be in the closed interval [0.0, 1.0]."""
273
+ top_k: Optional[int] = None
274
+ """Decode using top-k sampling: consider the set of top_k most probable tokens.
275
+ Must be positive."""
276
+ n: Optional[int] = None
277
+ """Number of chat completions to generate for each prompt. Note that the API may
278
+ not return the full n completions if duplicates are generated."""
279
+ max_tokens: Optional[int] = None
280
+
281
+ max_retries: int = 1
282
+
283
+ @property
284
+ def _default_params(self) -> Dict[str, Any]:
285
+ """Get the default parameters for calling OpenAI API."""
286
+ set_model_value = self.model
287
+ if self.model_name is not None:
288
+ set_model_value = self.model_name
289
+ return {
290
+ "model": set_model_value,
291
+ "force_timeout": self.request_timeout,
292
+ "max_tokens": self.max_tokens,
293
+ "stream": self.streaming,
294
+ "n": self.n,
295
+ "temperature": self.temperature,
296
+ "custom_llm_provider": self.custom_llm_provider,
297
+ **self.model_kwargs,
298
+ }
299
+
300
+ @property
301
+ def _client_params(self) -> Dict[str, Any]:
302
+ """Get the parameters used for the openai client."""
303
+ set_model_value = self.model
304
+ if self.model_name is not None:
305
+ set_model_value = self.model_name
306
+ self.client.api_base = self.api_base
307
+ self.client.api_key = self.api_key
308
+ for named_api_key in [
309
+ "openai_api_key",
310
+ "azure_api_key",
311
+ "anthropic_api_key",
312
+ "replicate_api_key",
313
+ "cohere_api_key",
314
+ "openrouter_api_key",
315
+ ]:
316
+ if api_key_value := getattr(self, named_api_key):
317
+ setattr(
318
+ self.client,
319
+ named_api_key.replace("_api_key", "_key"),
320
+ api_key_value,
321
+ )
322
+ self.client.organization = self.organization
323
+ creds: Dict[str, Any] = {
324
+ "model": set_model_value,
325
+ "force_timeout": self.request_timeout,
326
+ "api_base": self.api_base,
327
+ }
328
+ return {**self._default_params, **creds}
329
+
330
+ def completion_with_retry(
331
+ self, run_manager: Optional[CallbackManagerForLLMRun] = None, **kwargs: Any
332
+ ) -> Any:
333
+ """Use tenacity to retry the completion call."""
334
+ retry_decorator = _create_retry_decorator(self, run_manager=run_manager)
335
+
336
+ @retry_decorator
337
+ def _completion_with_retry(**kwargs: Any) -> Any:
338
+ return self.client.completion(**kwargs)
339
+
340
+ return _completion_with_retry(**kwargs)
341
+
342
+ @pre_init
343
+ def validate_environment(cls, values: Dict) -> Dict:
344
+ """Validate api key, python package exists, temperature, top_p, and top_k."""
345
+ try:
346
+ import litellm
347
+ except ImportError:
348
+ raise ChatLiteLLMException(
349
+ "Could not import litellm python package. "
350
+ "Please install it with `pip install litellm`"
351
+ )
352
+
353
+ values["openai_api_key"] = get_from_dict_or_env(
354
+ values, "openai_api_key", "OPENAI_API_KEY", default=""
355
+ )
356
+ values["azure_api_key"] = get_from_dict_or_env(
357
+ values, "azure_api_key", "AZURE_API_KEY", default=""
358
+ )
359
+ values["anthropic_api_key"] = get_from_dict_or_env(
360
+ values, "anthropic_api_key", "ANTHROPIC_API_KEY", default=""
361
+ )
362
+ values["replicate_api_key"] = get_from_dict_or_env(
363
+ values, "replicate_api_key", "REPLICATE_API_KEY", default=""
364
+ )
365
+ values["openrouter_api_key"] = get_from_dict_or_env(
366
+ values, "openrouter_api_key", "OPENROUTER_API_KEY", default=""
367
+ )
368
+ values["cohere_api_key"] = get_from_dict_or_env(
369
+ values, "cohere_api_key", "COHERE_API_KEY", default=""
370
+ )
371
+ values["huggingface_api_key"] = get_from_dict_or_env(
372
+ values, "huggingface_api_key", "HUGGINGFACE_API_KEY", default=""
373
+ )
374
+ values["together_ai_api_key"] = get_from_dict_or_env(
375
+ values, "together_ai_api_key", "TOGETHERAI_API_KEY", default=""
376
+ )
377
+ values["client"] = litellm
378
+
379
+ if values["temperature"] is not None and not 0 <= values["temperature"] <= 1:
380
+ raise ValueError("temperature must be in the range [0.0, 1.0]")
381
+
382
+ if values["top_p"] is not None and not 0 <= values["top_p"] <= 1:
383
+ raise ValueError("top_p must be in the range [0.0, 1.0]")
384
+
385
+ if values["top_k"] is not None and values["top_k"] <= 0:
386
+ raise ValueError("top_k must be positive")
387
+
388
+ return values
389
+
390
+ def _generate(
391
+ self,
392
+ messages: List[BaseMessage],
393
+ stop: Optional[List[str]] = None,
394
+ run_manager: Optional[CallbackManagerForLLMRun] = None,
395
+ stream: Optional[bool] = None,
396
+ **kwargs: Any,
397
+ ) -> ChatResult:
398
+ should_stream = stream if stream is not None else self.streaming
399
+ if should_stream:
400
+ stream_iter = self._stream(
401
+ messages, stop=stop, run_manager=run_manager, **kwargs
402
+ )
403
+ return generate_from_stream(stream_iter)
404
+
405
+ message_dicts, params = self._create_message_dicts(messages, stop)
406
+ params = {**params, **kwargs}
407
+ response = self.completion_with_retry(
408
+ messages=message_dicts, run_manager=run_manager, **params
409
+ )
410
+ return self._create_chat_result(response)
411
+
412
+ def _create_chat_result(self, response: Mapping[str, Any]) -> ChatResult:
413
+ generations = []
414
+ token_usage = response.get("usage", {})
415
+ for res in response["choices"]:
416
+ message = _convert_dict_to_message(res["message"])
417
+ if isinstance(message, AIMessage):
418
+ message.response_metadata = {
419
+ "model_name": self.model_name or self.model
420
+ }
421
+ message.usage_metadata = _create_usage_metadata(token_usage)
422
+ gen = ChatGeneration(
423
+ message=message,
424
+ generation_info=dict(finish_reason=res.get("finish_reason")),
425
+ )
426
+ generations.append(gen)
427
+ set_model_value = self.model
428
+ if self.model_name is not None:
429
+ set_model_value = self.model_name
430
+ llm_output = {"token_usage": token_usage, "model": set_model_value}
431
+ return ChatResult(generations=generations, llm_output=llm_output)
432
+
433
+ def _create_message_dicts(
434
+ self, messages: List[BaseMessage], stop: Optional[List[str]]
435
+ ) -> Tuple[List[Dict[str, Any]], Dict[str, Any]]:
436
+ params = self._client_params
437
+ if stop is not None:
438
+ if "stop" in params:
439
+ raise ValueError("`stop` found in both the input and default params.")
440
+ params["stop"] = stop
441
+ message_dicts = [_convert_message_to_dict(m) for m in messages]
442
+ return message_dicts, params
443
+
444
+ def _stream(
445
+ self,
446
+ messages: List[BaseMessage],
447
+ stop: Optional[List[str]] = None,
448
+ run_manager: Optional[CallbackManagerForLLMRun] = None,
449
+ **kwargs: Any,
450
+ ) -> Iterator[ChatGenerationChunk]:
451
+ message_dicts, params = self._create_message_dicts(messages, stop)
452
+ params = {**params, **kwargs, "stream": True}
453
+
454
+ default_chunk_class = AIMessageChunk
455
+ for chunk in self.completion_with_retry(
456
+ messages=message_dicts, run_manager=run_manager, **params
457
+ ):
458
+ if not isinstance(chunk, dict):
459
+ chunk = chunk.model_dump()
460
+ if len(chunk["choices"]) == 0:
461
+ continue
462
+ delta = chunk["choices"][0]["delta"]
463
+ chunk = _convert_delta_to_message_chunk(delta, default_chunk_class)
464
+ default_chunk_class = chunk.__class__
465
+ cg_chunk = ChatGenerationChunk(message=chunk)
466
+ if run_manager:
467
+ run_manager.on_llm_new_token(chunk.content, chunk=cg_chunk)
468
+ yield cg_chunk
469
+
470
+ async def _astream(
471
+ self,
472
+ messages: List[BaseMessage],
473
+ stop: Optional[List[str]] = None,
474
+ run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
475
+ **kwargs: Any,
476
+ ) -> AsyncIterator[ChatGenerationChunk]:
477
+ message_dicts, params = self._create_message_dicts(messages, stop)
478
+ params = {**params, **kwargs, "stream": True}
479
+
480
+ default_chunk_class = AIMessageChunk
481
+ async for chunk in await acompletion_with_retry(
482
+ self, messages=message_dicts, run_manager=run_manager, **params
483
+ ):
484
+ if not isinstance(chunk, dict):
485
+ chunk = chunk.model_dump()
486
+ if len(chunk["choices"]) == 0:
487
+ continue
488
+ delta = chunk["choices"][0]["delta"]
489
+ chunk = _convert_delta_to_message_chunk(delta, default_chunk_class)
490
+ default_chunk_class = chunk.__class__
491
+ cg_chunk = ChatGenerationChunk(message=chunk)
492
+ if run_manager:
493
+ await run_manager.on_llm_new_token(chunk.content, chunk=cg_chunk)
494
+ yield cg_chunk
495
+
496
+ async def _agenerate(
497
+ self,
498
+ messages: List[BaseMessage],
499
+ stop: Optional[List[str]] = None,
500
+ run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
501
+ stream: Optional[bool] = None,
502
+ **kwargs: Any,
503
+ ) -> ChatResult:
504
+ should_stream = stream if stream is not None else self.streaming
505
+ if should_stream:
506
+ stream_iter = self._astream(
507
+ messages=messages, stop=stop, run_manager=run_manager, **kwargs
508
+ )
509
+ return await agenerate_from_stream(stream_iter)
510
+
511
+ message_dicts, params = self._create_message_dicts(messages, stop)
512
+ params = {**params, **kwargs}
513
+ response = await acompletion_with_retry(
514
+ self, messages=message_dicts, run_manager=run_manager, **params
515
+ )
516
+ return self._create_chat_result(response)
517
+
518
+ def bind_tools(
519
+ self,
520
+ tools: Sequence[Union[Dict[str, Any], Type[BaseModel], Callable, BaseTool]],
521
+ tool_choice: Optional[
522
+ Union[dict, str, Literal["auto", "none", "required", "any"], bool]
523
+ ] = None,
524
+ **kwargs: Any,
525
+ ) -> Runnable[LanguageModelInput, BaseMessage]:
526
+ """Bind tool-like objects to this chat model.
527
+
528
+ LiteLLM expects tools argument in OpenAI format.
529
+
530
+ Args:
531
+ tools: A list of tool definitions to bind to this chat model.
532
+ Can be a dictionary, pydantic model, callable, or BaseTool. Pydantic
533
+ models, callables, and BaseTools will be automatically converted to
534
+ their schema dictionary representation.
535
+ tool_choice: Which tool to require the model to call. Options are:
536
+ - str of the form ``"<<tool_name>>"``: calls <<tool_name>> tool.
537
+ - ``"auto"``:
538
+ automatically selects a tool (including no tool).
539
+ - ``"none"``:
540
+ does not call a tool.
541
+ - ``"any"`` or ``"required"`` or ``True``:
542
+ forces least one tool to be called.
543
+ - dict of the form:
544
+ ``{"type": "function", "function": {"name": <<tool_name>>}}``
545
+ - ``False`` or ``None``: no effect
546
+ **kwargs: Any additional parameters to pass to the
547
+ :class:`~langchain.runnable.Runnable` constructor.
548
+ """
549
+
550
+ formatted_tools = [convert_to_openai_tool(tool) for tool in tools]
551
+
552
+ # In case of openai if tool_choice is `any` or if bool has been provided we
553
+ # change it to `required` as that is suppored by openai.
554
+ if (
555
+ (self.model is not None and "azure" in self.model)
556
+ or (self.model_name is not None and "azure" in self.model_name)
557
+ or (self.model is not None and self.model in _OPENAI_MODELS)
558
+ or (self.model_name is not None and self.model_name in _OPENAI_MODELS)
559
+ ) and (tool_choice == "any" or isinstance(tool_choice, bool)):
560
+ tool_choice = "required"
561
+ # If tool_choice is bool apart from openai we make it `any`
562
+ elif isinstance(tool_choice, bool):
563
+ tool_choice = "any"
564
+ elif isinstance(tool_choice, dict):
565
+ tool_names = [
566
+ formatted_tool["function"]["name"] for formatted_tool in formatted_tools
567
+ ]
568
+ if not any(
569
+ tool_name == tool_choice["function"]["name"] for tool_name in tool_names
570
+ ):
571
+ raise ValueError(
572
+ f"Tool choice {tool_choice} was specified, but the only "
573
+ f"provided tools were {tool_names}."
574
+ )
575
+ return super().bind(tools=formatted_tools, tool_choice=tool_choice, **kwargs)
576
+
577
+ @property
578
+ def _identifying_params(self) -> Dict[str, Any]:
579
+ """Get the identifying parameters."""
580
+ set_model_value = self.model
581
+ if self.model_name is not None:
582
+ set_model_value = self.model_name
583
+ return {
584
+ "model": set_model_value,
585
+ "temperature": self.temperature,
586
+ "top_p": self.top_p,
587
+ "top_k": self.top_k,
588
+ "n": self.n,
589
+ }
590
+
591
+ @property
592
+ def _llm_type(self) -> str:
593
+ return "litellm-chat"
594
+
595
+
596
+ def _create_usage_metadata(token_usage: Mapping[str, Any]) -> UsageMetadata:
597
+ input_tokens = token_usage.get("prompt_tokens", 0)
598
+ output_tokens = token_usage.get("completion_tokens", 0)
599
+ return UsageMetadata(
600
+ input_tokens=input_tokens,
601
+ output_tokens=output_tokens,
602
+ total_tokens=input_tokens + output_tokens,
603
+ )
File without changes
@@ -0,0 +1,99 @@
1
+ [build-system]
2
+ requires = ["poetry-core>=1.0.0"]
3
+ build-backend = "poetry.core.masonry.api"
4
+
5
+ [tool.poetry]
6
+ name = "langchain-litellm"
7
+ version = "0.1.0"
8
+ description = "An integration package connecting Litellm and LangChain"
9
+ authors = []
10
+ readme = "README.md"
11
+ repository = "https://github.com/langchain-ai/langchain"
12
+ license = "MIT"
13
+
14
+ [tool.mypy]
15
+ disallow_untyped_defs = "True"
16
+
17
+ [tool.poetry.urls]
18
+ "Source Code" = "https://github.com/langchain-ai/langchain/tree/master/libs/partners/litellm"
19
+ "Release Notes" = "https://github.com/langchain-ai/langchain/releases?q=tag%3A%22litellm%3D%3D0%22&expanded=true"
20
+
21
+ [tool.poetry.dependencies]
22
+ python = ">=3.9,<4.0"
23
+ langchain-core = "^0.3.15"
24
+ litellm = "^1.65.1"
25
+
26
+ [tool.ruff.lint]
27
+ select = ["E", "F", "I", "T201"]
28
+
29
+ [tool.coverage.run]
30
+ omit = ["tests/*"]
31
+
32
+ [tool.pytest.ini_options]
33
+ addopts = "--strict-markers --strict-config --durations=5"
34
+ markers = [
35
+ "compile: mark placeholder test used to compile integration tests without running them",
36
+ ]
37
+ asyncio_mode = "auto"
38
+
39
+ [tool.poetry.group.test]
40
+ optional = true
41
+
42
+ [tool.poetry.group.codespell]
43
+ optional = true
44
+
45
+ [tool.poetry.group.test_integration]
46
+ optional = true
47
+
48
+ [tool.poetry.group.lint]
49
+ optional = true
50
+
51
+ [tool.poetry.group.dev]
52
+ optional = true
53
+
54
+ [tool.poetry.group.dev.dependencies]
55
+ ipykernel = "^6.29.5"
56
+
57
+ [tool.poetry.group.test.dependencies]
58
+ pytest = ">=7.4.4,<8.0.0"
59
+ pytest-asyncio = ">=0.20.3,<1.0.0"
60
+ pytest-socket = ">=0.6.0,<1.0.0"
61
+ pytest-watcher = ">=0.2.6,<1.0.0"
62
+ langchain-tests = "0.3.17"
63
+ pytest-cov = ">=4.1.0,<5.0.0"
64
+ pytest-dotenv = ">=0.5.2,<1.0.0"
65
+ duckdb-engine = ">=0.13.6,<1.0.0"
66
+ freezegun = ">=1.2.2,<2.0.0"
67
+ responses = ">=0.22.0,<1.0.0"
68
+ lark = ">=1.1.5,<2.0.0"
69
+ pandas = ">=2.0.0,<3.0.0"
70
+ pytest-mock = ">=3.10.0,<4.0.0"
71
+ syrupy = ">=4.0.2,<5.0.0"
72
+ requests-mock = ">=1.11.0,<2.0.0"
73
+ pytest-xdist = ">=3.6.1,<4.0.0"
74
+ blockbuster = ">=1.5.18,<1.6"
75
+ cffi = {markers = "python_version >= \"3.10\"", version = "^1.17.1"}
76
+ langchain-core = "^0.3.49"
77
+ langchain = "^0.3.22"
78
+ toml = ">=0.10.2"
79
+
80
+ [tool.poetry.group.codespell.dependencies]
81
+ codespell = "^2.2.6"
82
+
83
+ [tool.poetry.group.test_integration.dependencies]
84
+
85
+ [tool.poetry.group.lint.dependencies]
86
+ ruff = "^0.5"
87
+
88
+ [tool.poetry.group.typing.dependencies]
89
+ mypy = ">=1.12,<2.0"
90
+ types-pyyaml = ">=6.0.12.2,<7.0.0.0"
91
+ types-requests = ">=2.28.11.5,<3.0.0.0"
92
+ types-toml = ">=0.10.8.1,<1.0.0.0"
93
+ types-pytz = ">=2023.3.0.0,<2024.0.0.0"
94
+ types-chardet = ">=5.0.4.6,<6.0.0.0"
95
+ types-redis = ">=4.3.21.6,<5.0.0.0"
96
+ mypy-protobuf = ">=3.0.0,<4.0.0"
97
+ langchain-core = "^0.3.49"
98
+ langchain-text-splitters = "^0.3.7"
99
+ langchain = "^0.3.22"