fix: Do a better job of mapping Langchain to LiteLLM (#5233)
This commit is contained in:
parent
0531084e35
commit
a17335e802
1 changed files with 15 additions and 6 deletions
|
|
@ -1,6 +1,7 @@
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
|
|
||||||
|
import litellm
|
||||||
from crewai import LLM, Agent, Crew, Process, Task
|
from crewai import LLM, Agent, Crew, Process, Task
|
||||||
from crewai.task import TaskOutput
|
from crewai.task import TaskOutput
|
||||||
from crewai.tools.base_tool import Tool
|
from crewai.tools.base_tool import Tool
|
||||||
|
|
@ -70,9 +71,12 @@ def convert_llm(llm: Any, excluded_keys=None) -> LLM:
|
||||||
msg = "Could not find model name in the LLM object"
|
msg = "Could not find model name in the LLM object"
|
||||||
raise ValueError(msg)
|
raise ValueError(msg)
|
||||||
|
|
||||||
# Normalize Ollama with prefix TODO: Handle all litellm supported models
|
# Normalize to the LLM model name
|
||||||
if llm.dict().get("_type") == "chat-ollama":
|
# Remove langchain_ prefix if present
|
||||||
model_name = f"ollama/{model_name}"
|
provider = llm.get_lc_namespace()[0]
|
||||||
|
if provider.startswith("langchain_"):
|
||||||
|
provider = provider[10:]
|
||||||
|
model_name = f"{provider}/{model_name}"
|
||||||
|
|
||||||
# Retrieve the API Key from the LLM
|
# Retrieve the API Key from the LLM
|
||||||
if excluded_keys is None:
|
if excluded_keys is None:
|
||||||
|
|
@ -188,8 +192,13 @@ class BaseCrewComponent(Component):
|
||||||
return step_callback
|
return step_callback
|
||||||
|
|
||||||
async def build_output(self) -> Message:
|
async def build_output(self) -> Message:
|
||||||
crew = self.build_crew()
|
try:
|
||||||
result = await crew.kickoff_async()
|
crew = self.build_crew()
|
||||||
message = Message(text=result.raw, sender=MESSAGE_SENDER_AI)
|
result = await crew.kickoff_async()
|
||||||
|
message = Message(text=result.raw, sender=MESSAGE_SENDER_AI)
|
||||||
|
except litellm.exceptions.BadRequestError as e:
|
||||||
|
raise ValueError(e) from e
|
||||||
|
|
||||||
self.status = message
|
self.status = message
|
||||||
|
|
||||||
return message
|
return message
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue