LiteLLM integration (#419)

TerminallyLazy committed Jun 10, 2025 at 14:21 UTC e93cb72b096323b0188273c4b56423fb95811665
5 files changed +245 -420
agent.py
+21 -8
@@ -5,7 +5,7 @@ nest_asyncio.apply()
5 from collections import OrderedDict
6 from dataclasses import dataclass, field
7 from datetime import datetime, timezone
8 -from typing import Any, Awaitable, Coroutine, Dict
8 +from typing import Any, Awaitable, Coroutine, Optional, Dict, TypedDict, cast
9 from enum import Enum
10 import uuid
11 import models
@@ -45,9 +45,9 @@ class AgentContext:
45 log: Log.Log | None = None,
46 paused: bool = False,
47 streaming_agent: "Agent|None" = None,
48 - created_at: datetime | None = None,
48 + created_at: "datetime|None" = None,
49 type: AgentContextType = AgentContextType.USER,
50 - last_message: datetime | None = None,
50 + last_message: datetime |None = None,
51 ):
52 # build context
53 self.id = id or str(uuid.uuid4())
@@ -111,6 +111,7 @@ class AgentContext:
111 "type": self.type.value,
112 }
113
114 +
115 @staticmethod
116 def log_to_all(
117 type: Log.Type,
@@ -611,10 +612,14 @@ class Agent:
612 self.config.utility_model, prompt.format(), background
613 )
614
614 - async for chunk in (prompt | model).astream({}):
615 + # async for chunk in (prompt | model).astream({}):
616 + # await self.handle_intervention() # wait for intervention and handle it, if paused
617 + # format prompt to messages and stream from model directly
618 + model_stream = cast(Any, model)
619 + messages = prompt.format_messages()
620 + async for chunk in model_stream.astream(messages):
621 await self.handle_intervention() # wait for intervention and handle it, if paused
616 -
617 - content = models.parse_chunk(chunk)
622 + content = chunk if isinstance(chunk, str) else str(chunk)
623 limiter.add(output=tokens.approximate_tokens(content))
624 response += content
625
@@ -636,13 +641,20 @@ class Agent:
641 # rate limiter
642 limiter = await self.rate_limiter(self.config.chat_model, prompt.format())
643
639 - async for chunk in (prompt | model).astream({}):
644 + # async for chunk in (prompt | model).astream({}):
645 + # await self.handle_intervention() # wait for intervention and handle it, if paused
646 + # format prompt to messages and stream from model directly
647 + messages = prompt.format_messages()
648 + model_stream = cast(Any, model)
649 + async for chunk in model_stream.astream(messages):
650 await self.handle_intervention() # wait for intervention and handle it, if paused
651
652 content = models.parse_chunk(chunk)
653 + content = chunk if isinstance(chunk, str) else str(chunk)
654 limiter.add(output=tokens.approximate_tokens(content))
655 response += content
656
657 +
658 if callback:
659 await callback(content, response)
660
@@ -713,7 +725,8 @@ class Agent:
725 if ":" in raw_tool_name:
726 tool_name, tool_method = raw_tool_name.split(":", 1)
727
716 - tool = None # Initialize tool to None
728 + tool = None
729 + # tool = self.get_tool(name=tool_name, method=tool_method, args=tool_args, message=msg)
730
731 # Try getting tool from MCP first
732 try:
models.py
+210 -396
@@ -1,47 +1,24 @@
1 -from enum import Enum
2 -import os
3 -from typing import Any
4 -from langchain_openai import (
5 - ChatOpenAI,
6 - OpenAI,
7 - OpenAIEmbeddings,
8 - AzureChatOpenAI,
9 - AzureOpenAIEmbeddings,
10 - AzureOpenAI,
11 -)
12 -from langchain_community.llms.ollama import Ollama
13 -from langchain_ollama import ChatOllama
14 -from langchain_community.embeddings import OllamaEmbeddings
15 -from langchain_anthropic import ChatAnthropic
16 -from langchain_groq import ChatGroq
17 -from langchain_huggingface import (
18 - HuggingFaceEmbeddings,
19 - ChatHuggingFace,
20 - HuggingFaceEndpoint,
21 -)
22 -from langchain_google_genai import (
23 - ChatGoogleGenerativeAI,
24 - HarmBlockThreshold,
25 - HarmCategory,
26 - embeddings as google_embeddings,
27 -)
28 -from langchain_mistralai import ChatMistralAI
29 -
30 -# from pydantic.v1.types import SecretStr
31 -from python.helpers import dotenv, runtime
32 -from python.helpers.dotenv import load_dotenv
33 -from python.helpers.rate_limiter import RateLimiter
34 -
35 -# environment variables
36 -load_dotenv()
37 -
38 -
39 -class ModelType(Enum):
40 - CHAT = "Chat"
41 - EMBEDDING = "Embedding"
42 -
43 -
44 -class ModelProvider(Enum):
1 +from __future__ import annotations
2 +from enum import Enum
3 +import os
4 +from typing import Any,AsyncIterator, Iterable, List, Dict, Union, cast
5 +
6 +import litellm
7 +from langchain_core.embeddings import Embeddings
8 +from langchain_core.runnables import RunnableLambda
9 +from langchain_core.messages import BaseMessage
10 +from python.helpers import dotenv, runtime
11 +from python.helpers.dotenv import load_dotenv
12 +from python.helpers.rate_limiter import RateLimiter
13 +
14 +# environment variables
15 +load_dotenv()
16 +
17 +class ModelType(Enum):
18 + CHAT = "Chat"
19 + EMBEDDING = "Embedding"
20 +
21 +class ModelProvider(Enum):
22 ANTHROPIC = "Anthropic"
23 CHUTES = "Chutes"
24 DEEPSEEK = "DeepSeek"
@@ -61,24 +38,199 @@ class ModelProvider(Enum):
38 rate_limiters: dict[str, RateLimiter] = {}
39
40
64 -# Utility function to get API keys from environment variables
65 -def get_api_key(service):
66 - return (
67 - dotenv.get_dotenv_value(f"API_KEY_{service.upper()}")
41 +# ---------------------------------------------------------------------------
42 +# helpers
43 +# ---------------------------------------------------------------------------
44 +
45 +# Utility function to get API keys from environment variables
46 +def get_api_key(service):
47 + return (
48 + dotenv.get_dotenv_value(f"API_KEY_{service.upper()}")
49 or dotenv.get_dotenv_value(f"{service.upper()}_API_KEY")
69 - or dotenv.get_dotenv_value(
70 - f"{service.upper()}_API_TOKEN"
71 - ) # Added for CHUTES_API_TOKEN
72 - or "None"
50 + # or dotenv.get_dotenv_value(f"{service.upper()}_API_TOKEN")
51 + or "None"
52 )
53
54
76 -def get_model(type: ModelType, provider: ModelProvider, name: str, **kwargs):
77 - fnc_name = f"get_{provider.name.lower()}_{type.name.lower()}" # function name of model getter
78 - model = globals()[fnc_name](name, **kwargs) # call function by name
79 - return model
80 -
81 -
55 +BASE_URL_ENV = {
56 + ModelProvider.OLLAMA: "OLLAMA_BASE_URL",
57 + ModelProvider.LMSTUDIO: "LM_STUDIO_BASE_URL",
58 + ModelProvider.ANTHROPIC: "ANTHROPIC_BASE_URL",
59 + ModelProvider.DEEPSEEK: "DEEPSEEK_BASE_URL",
60 + ModelProvider.OPENROUTER: "OPEN_ROUTER_BASE_URL",
61 + ModelProvider.SAMBANOVA: "SAMBANOVA_BASE_URL",
62 + ModelProvider.CHUTES: "CHUTES_BASE_URL",
63 +}
64 +
65 +def get_base_url(provider: ModelProvider) -> str | None:
66 + if provider == ModelProvider.OLLAMA:
67 + return dotenv.get_dotenv_value("OLLAMA_BASE_URL") or f"http://{runtime.get_local_url()}:11434"
68 + if provider == ModelProvider.LMSTUDIO:
69 + return dotenv.get_dotenv_value("LM_STUDIO_BASE_URL") or f"http://{runtime.get_local_url()}:1234/v1"
70 + env = BASE_URL_ENV.get(provider)
71 + if env:
72 + return dotenv.get_dotenv_value(env)
73 + if provider == ModelProvider.OPENAI_AZURE:
74 + return dotenv.get_dotenv_value("OPENAI_AZURE_ENDPOINT")
75 + return None
76 +
77 +def parse_chunk(chunk: Any) -> str:
78 + """Parse a streaming chunk from LiteLLM."""
79 + if isinstance(chunk, str):
80 + return chunk
81 +
82 + # Handle LiteLLM ModelResponseStream objects
83 + if hasattr(chunk, "choices") and chunk.choices:
84 + choice = chunk.choices[0]
85 + if hasattr(choice, "delta") and choice.delta:
86 + delta = choice.delta
87 + if hasattr(delta, "content") and delta.content:
88 + return str(delta.content)
89 +
90 + # Some SDKs return an object with a `content` attribute rather than a dict.
91 + if hasattr(chunk, "content"):
92 + return str(chunk.content)
93 +
94 + # Handle dict format
95 + if isinstance(chunk, dict):
96 + delta = (
97 + chunk.get("choices", [{}])[0]
98 + .get("delta", {})
99 + .get("content")
100 + )
101 + if delta:
102 + return str(delta)
103 + return str(chunk)
104 +
105 +# ---------------------------------------------------------------------------
106 +# LiteLLM wrappers
107 +# ---------------------------------------------------------------------------
108 +
109 +class LiteLLMEmbeddings(Embeddings):
110 + """LangChain embeddings wrapper around LiteLLM."""
111 +
112 + def __init__(self, model: str, *, api_key: str | None = None, api_base: str | None = None, **kwargs: Any) -> None:
113 + self.model = model
114 + self.api_key = api_key
115 + self.api_base = api_base
116 + self.kwargs = kwargs
117 +
118 + def _embedding_args(self) -> Dict[str, Any]:
119 + args = {"model": self.model, **self.kwargs}
120 + if self.api_key:
121 + args["api_key"] = self.api_key
122 + if self.api_base:
123 + args["api_base"] = self.api_base
124 + return args
125 +
126 + def embed_documents(self, texts: List[str]) -> List[List[float]]:
127 + resp = litellm.embedding(input=texts, **self._embedding_args()) # type: ignore
128 + return [d["embedding"] for d in resp["data"]] # type: ignore[index]
129 +
130 + def embed_query(self, text: str) -> List[float]:
131 + return self.embed_documents([text])[0]
132 +
133 +# ---------------------------------------------------------------------------
134 +
135 +
136 +def _to_litellm_messages(messages: Iterable[BaseMessage]) -> List[Dict[str, str]]:
137 + llm_messages = []
138 + for m in messages:
139 + # Convert LangChain message types to LiteLLM/OpenAI format
140 + if m.type == "ai":
141 + role = "assistant"
142 + elif m.type == "human":
143 + role = "user" # LiteLLM expects "user" not "human"
144 + elif m.type == "system":
145 + role = "system"
146 + else:
147 + role = "user" # Default fallback
148 +
149 + llm_messages.append({"role": role, "content": m.content})
150 + return llm_messages
151 +
152 +
153 +def _convert_numeric_params(kwargs: dict) -> dict:
154 + """Convert string numeric parameters to proper types for LiteLLM"""
155 + converted = kwargs.copy()
156 +
157 + # Parameters that should be converted to float
158 + float_params = ['temperature', 'top_p', 'frequency_penalty', 'presence_penalty']
159 + # Parameters that should be converted to int
160 + int_params = ['max_tokens', 'n', 'seed', 'timeout']
161 +
162 + for param in float_params:
163 + if param in converted and isinstance(converted[param], str):
164 + try:
165 + converted[param] = float(converted[param])
166 + except (ValueError, TypeError):
167 + pass # Keep original value if conversion fails
168 +
169 + for param in int_params:
170 + if param in converted and isinstance(converted[param], str):
171 + try:
172 + converted[param] = int(converted[param])
173 + except (ValueError, TypeError):
174 + pass # Keep original value if conversion fails
175 +
176 + return converted
177 +
178 +
179 +def get_chat_model(provider: ModelProvider, name: str, **kwargs: Any) -> RunnableLambda:
180 + # Convert string numeric parameters to proper types
181 + kwargs = _convert_numeric_params(kwargs)
182 +
183 + api_key = kwargs.pop("api_key", None) or get_api_key(provider.name)
184 + api_base = kwargs.pop("api_base", None) or get_base_url(provider)
185 +
186 + if provider == ModelProvider.OPENAI_AZURE:
187 + kwargs.setdefault("custom_llm_provider", "azure")
188 + version = dotenv.get_dotenv_value("OPENAI_API_VERSION")
189 + if version:
190 + kwargs.setdefault("api_version", version)
191 +
192 + async def _chat(messages: List[BaseMessage]) -> AsyncIterator[str]:
193 + llm_messages = _to_litellm_messages(messages)
194 + # Await the coroutine returned by `litellm.acompletion` to obtain the
195 + # stream wrapper that implements `__aiter__`, then iterate over it.
196 + stream_wrapper = await litellm.acompletion(
197 + model=name,
198 + messages=llm_messages,
199 + stream=True,
200 + api_key=api_key,
201 + api_base=api_base,
202 + **kwargs,
203 + )
204 + async for chunk in cast(AsyncIterator[Any], stream_wrapper):
205 + text = parse_chunk(chunk)
206 + if text:
207 + yield text
208 +
209 + return RunnableLambda(lambda msgs: _chat(msgs), afunc=_chat, name=f"litellm_chat_{name}")
210 +
211 +
212 +def get_embedding_model(provider: ModelProvider, name: str, **kwargs: Any) -> LiteLLMEmbeddings:
213 + api_key = kwargs.pop("api_key", None) or get_api_key(provider.name)
214 + api_base = kwargs.pop("api_base", None) or get_base_url(provider)
215 + if provider == ModelProvider.OPENAI_AZURE:
216 + kwargs.setdefault("custom_llm_provider", "azure")
217 + version = dotenv.get_dotenv_value("OPENAI_API_VERSION")
218 + if version:
219 + kwargs.setdefault("api_version", version)
220 + return LiteLLMEmbeddings(model=name, api_key=api_key, api_base=api_base, **kwargs)
221 +
222 +# ---------------------------------------------------------------------------
223 +# public API
224 +# ---------------------------------------------------------------------------
225 +
226 +def get_model(type: ModelType, provider: ModelProvider, name: str, **kwargs: Any) -> Union[RunnableLambda, LiteLLMEmbeddings]:
227 + if type == ModelType.CHAT:
228 + return get_chat_model(provider, name, **kwargs)
229 + elif type == ModelType.EMBEDDING:
230 + return get_embedding_model(provider, name, **kwargs)
231 + else:
232 + raise ValueError(f"Unsupported model type: {type}")
233 +
234 def get_rate_limiter(
235 provider: ModelProvider, name: str, requests: int, input: int, output: int
236 ) -> RateLimiter:
@@ -92,342 +244,4 @@ def get_rate_limiter(
244 return limiter
245
246
95 -def parse_chunk(chunk: Any):
96 - if isinstance(chunk, str):
97 - content = chunk
98 - elif hasattr(chunk, "content"):
99 - content = str(chunk.content)
100 - else:
101 - content = str(chunk)
102 - return content
103 -
104 -
105 -# Ollama models
106 -def get_ollama_base_url():
107 - return (
108 - dotenv.get_dotenv_value("OLLAMA_BASE_URL")
109 - or f"http://{runtime.get_local_url()}:11434"
110 - )
111 -
112 -
113 -def get_ollama_chat(
114 - model_name: str,
115 - base_url=None,
116 - num_ctx=8192,
117 - **kwargs,
118 -):
119 - if not base_url:
120 - base_url = get_ollama_base_url()
121 - return ChatOllama(
122 - model=model_name,
123 - base_url=base_url,
124 - num_ctx=num_ctx,
125 - **kwargs,
126 - )
127 -
128 -
129 -def get_ollama_embedding(
130 - model_name: str,
131 - base_url=None,
132 - num_ctx=8192,
133 - **kwargs,
134 -):
135 - if not base_url:
136 - base_url = get_ollama_base_url()
137 - return OllamaEmbeddings(
138 - model=model_name, base_url=base_url, num_ctx=num_ctx, **kwargs
139 - )
247
141 -
142 -# HuggingFace models
143 -def get_huggingface_chat(
144 - model_name: str,
145 - api_key=None,
146 - **kwargs,
147 -):
148 - # different naming convention here
149 - if not api_key:
150 - api_key = get_api_key("huggingface") or os.environ["HUGGINGFACEHUB_API_TOKEN"]
151 -
152 - # Initialize the HuggingFaceEndpoint with the specified model and parameters
153 - llm = HuggingFaceEndpoint(
154 - repo_id=model_name,
155 - task="text-generation",
156 - do_sample=True,
157 - **kwargs,
158 - )
159 -
160 - # Initialize the ChatHuggingFace with the configured llm
161 - return ChatHuggingFace(llm=llm)
162 -
163 -
164 -def get_huggingface_embedding(model_name: str, **kwargs):
165 - return HuggingFaceEmbeddings(model_name=model_name, **kwargs)
166 -
167 -
168 -# LM Studio and other OpenAI compatible interfaces
169 -def get_lmstudio_base_url():
170 - return (
171 - dotenv.get_dotenv_value("LM_STUDIO_BASE_URL")
172 - or f"http://{runtime.get_local_url()}:1234/v1"
173 - )
174 -
175 -
176 -def get_lmstudio_chat(
177 - model_name: str,
178 - base_url=None,
179 - **kwargs,
180 -):
181 - if not base_url:
182 - base_url = get_lmstudio_base_url()
183 - return ChatOpenAI(model_name=model_name, base_url=base_url, api_key="none", **kwargs) # type: ignore
184 -
185 -
186 -def get_lmstudio_embedding(
187 - model_name: str,
188 - base_url=None,
189 - **kwargs,
190 -):
191 - if not base_url:
192 - base_url = get_lmstudio_base_url()
193 - return OpenAIEmbeddings(model=model_name, api_key="none", base_url=base_url, check_embedding_ctx_length=False, **kwargs) # type: ignore
194 -
195 -
196 -# Anthropic models
197 -def get_anthropic_chat(
198 - model_name: str,
199 - api_key=None,
200 - base_url=None,
201 - **kwargs,
202 -):
203 - if not api_key:
204 - api_key = get_api_key("anthropic")
205 - if not base_url:
206 - base_url = (
207 - dotenv.get_dotenv_value("ANTHROPIC_BASE_URL") or "https://api.anthropic.com"
208 - )
209 - return ChatAnthropic(model_name=model_name, api_key=api_key, base_url=base_url, **kwargs) # type: ignore
210 -
211 -
212 -# right now anthropic does not have embedding models, but that might change
213 -def get_anthropic_embedding(
214 - model_name: str,
215 - api_key=None,
216 - **kwargs,
217 -):
218 - if not api_key:
219 - api_key = get_api_key("anthropic")
220 - return OpenAIEmbeddings(model=model_name, api_key=api_key, **kwargs) # type: ignore
221 -
222 -
223 -# OpenAI models
224 -def get_openai_chat(
225 - model_name: str,
226 - api_key=None,
227 - **kwargs,
228 -):
229 - if not api_key:
230 - api_key = get_api_key("openai")
231 - return ChatOpenAI(model_name=model_name, api_key=api_key, **kwargs) # type: ignore
232 -
233 -
234 -def get_openai_embedding(model_name: str, api_key=None, **kwargs):
235 - if not api_key:
236 - api_key = get_api_key("openai")
237 - return OpenAIEmbeddings(model=model_name, api_key=api_key, **kwargs) # type: ignore
238 -
239 -
240 -def get_openai_azure_chat(
241 - deployment_name: str,
242 - api_key=None,
243 - azure_endpoint=None,
244 - **kwargs,
245 -):
246 - if not api_key:
247 - api_key = get_api_key("openai_azure")
248 - if not azure_endpoint:
249 - azure_endpoint = dotenv.get_dotenv_value("OPENAI_AZURE_ENDPOINT")
250 - return AzureChatOpenAI(deployment_name=deployment_name, api_key=api_key, azure_endpoint=azure_endpoint, **kwargs) # type: ignore
251 -
252 -
253 -def get_openai_azure_embedding(
254 - deployment_name: str,
255 - api_key=None,
256 - azure_endpoint=None,
257 - **kwargs,
258 -):
259 - if not api_key:
260 - api_key = get_api_key("openai_azure")
261 - if not azure_endpoint:
262 - azure_endpoint = dotenv.get_dotenv_value("OPENAI_AZURE_ENDPOINT")
263 - return AzureOpenAIEmbeddings(deployment_name=deployment_name, api_key=api_key, azure_endpoint=azure_endpoint, **kwargs) # type: ignore
264 -
265 -
266 -# Google models
267 -def get_google_chat(
268 - model_name: str,
269 - api_key=None,
270 - **kwargs,
271 -):
272 - if not api_key:
273 - api_key = get_api_key("google")
274 - return ChatGoogleGenerativeAI(model=model_name, google_api_key=api_key, safety_settings={HarmCategory.HARM_CATEGORY_DANGEROUS_CONTENT: HarmBlockThreshold.BLOCK_NONE}, **kwargs) # type: ignore
275 -
276 -
277 -def get_google_embedding(
278 - model_name: str,
279 - api_key=None,
280 - **kwargs,
281 -):
282 - if not api_key:
283 - api_key = get_api_key("google")
284 - return google_embeddings.GoogleGenerativeAIEmbeddings(model=model_name, google_api_key=api_key, **kwargs) # type: ignore
285 -
286 -
287 -# Mistral models
288 -def get_mistralai_chat(
289 - model_name: str,
290 - api_key=None,
291 - **kwargs,
292 -):
293 - if not api_key:
294 - api_key = get_api_key("mistral")
295 - return ChatMistralAI(model=model_name, api_key=api_key, **kwargs) # type: ignore
296 -
297 -
298 -# Groq models
299 -def get_groq_chat(
300 - model_name: str,
301 - api_key=None,
302 - **kwargs,
303 -):
304 - if not api_key:
305 - api_key = get_api_key("groq")
306 - return ChatGroq(model_name=model_name, api_key=api_key, **kwargs) # type: ignore
307 -
308 -
309 -# DeepSeek models
310 -def get_deepseek_chat(
311 - model_name: str,
312 - api_key=None,
313 - base_url=None,
314 - **kwargs,
315 -):
316 - if not api_key:
317 - api_key = get_api_key("deepseek")
318 - if not base_url:
319 - base_url = (
320 - dotenv.get_dotenv_value("DEEPSEEK_BASE_URL") or "https://api.deepseek.com"
321 - )
322 -
323 - return ChatOpenAI(api_key=api_key, model=model_name, base_url=base_url, **kwargs) # type: ignore
324 -
325 -
326 -# OpenRouter models
327 -def get_openrouter_chat(
328 - model_name: str,
329 - api_key=None,
330 - base_url=None,
331 - **kwargs,
332 -):
333 - if not api_key:
334 - api_key = get_api_key("openrouter")
335 - if not base_url:
336 - base_url = (
337 - dotenv.get_dotenv_value("OPEN_ROUTER_BASE_URL")
338 - or "https://openrouter.ai/api/v1"
339 - )
340 - return ChatOpenAI(
341 - api_key=api_key, # type: ignore
342 - model=model_name,
343 - base_url=base_url,
344 - stream_usage=True,
345 - model_kwargs={
346 - "extra_headers": {
347 - "HTTP-Referer": "https://agent-zero.ai",
348 - "X-Title": "Agent Zero",
349 - }
350 - },
351 - **kwargs,
352 - )
353 -
354 -
355 -def get_openrouter_embedding(
356 - model_name: str,
357 - api_key=None,
358 - base_url=None,
359 - **kwargs,
360 -):
361 - if not api_key:
362 - api_key = get_api_key("openrouter")
363 - if not base_url:
364 - base_url = (
365 - dotenv.get_dotenv_value("OPEN_ROUTER_BASE_URL")
366 - or "https://openrouter.ai/api/v1"
367 - )
368 - return OpenAIEmbeddings(model=model_name, api_key=api_key, base_url=base_url, **kwargs) # type: ignore
369 -
370 -
371 -# Sambanova models
372 -def get_sambanova_chat(
373 - model_name: str,
374 - api_key=None,
375 - base_url=None,
376 - max_tokens=1024,
377 - **kwargs,
378 -):
379 - if not api_key:
380 - api_key = get_api_key("sambanova")
381 - if not base_url:
382 - base_url = (
383 - dotenv.get_dotenv_value("SAMBANOVA_BASE_URL")
384 - or "https://fast-api.snova.ai/v1"
385 - )
386 - return ChatOpenAI(api_key=api_key, model=model_name, base_url=base_url, max_tokens=max_tokens, **kwargs) # type: ignore
387 -
388 -
389 -# right now sambanova does not have embedding models, but that might change
390 -def get_sambanova_embedding(
391 - model_name: str,
392 - api_key=None,
393 - base_url=None,
394 - **kwargs,
395 -):
396 - if not api_key:
397 - api_key = get_api_key("sambanova")
398 - if not base_url:
399 - base_url = (
400 - dotenv.get_dotenv_value("SAMBANOVA_BASE_URL")
401 - or "https://fast-api.snova.ai/v1"
402 - )
403 - return OpenAIEmbeddings(model=model_name, api_key=api_key, base_url=base_url, **kwargs) # type: ignore
404 -
405 -
406 -# Other OpenAI compatible models
407 -def get_other_chat(
408 - model_name: str,
409 - api_key=None,
410 - base_url=None,
411 - **kwargs,
412 -):
413 - return ChatOpenAI(api_key=api_key, model=model_name, base_url=base_url, **kwargs) # type: ignore
414 -
415 -
416 -def get_other_embedding(model_name: str, api_key=None, base_url=None, **kwargs):
417 - return OpenAIEmbeddings(model=model_name, api_key=api_key, base_url=base_url, **kwargs) # type: ignore
418 -
419 -
420 -# Chutes models
421 -def get_chutes_chat(
422 - model_name: str,
423 - api_key=None,
424 - base_url=None,
425 - **kwargs,
426 -):
427 - if not api_key:
428 - api_key = get_api_key("chutes")
429 - if not base_url:
430 - base_url = (
431 - dotenv.get_dotenv_value("CHUTES_BASE_URL") or "https://llm.chutes.ai/v1"
432 - )
433 - return ChatOpenAI(api_key=api_key, model=model_name, base_url=base_url, **kwargs) # type: ignore
python/helpers/tool.py
+4 -3
@@ -33,11 +33,12 @@ class Tool:
33 PrintStyle().print()
34
35 async def after_execution(self, response: Response, **kwargs):
36 - text = response.message.strip()
36 + message = response.message or ""
37 + text = message.strip()
38 self.agent.hist_add_tool_result(self.name, text)
39 PrintStyle(font_color="#1B4F72", background_color="white", padding=True, bold=True).print(f"{self.agent.agent_name}: Response from tool '{self.name}'")
39 - PrintStyle(font_color="#85C1E9").print(response.message)
40 - self.log.update(content=response.message)
40 + PrintStyle(font_color="#85C1E9").print(message)
41 + self.log.update(content=message)
42
43 def get_log_object(self):
44 if self.method:
python/tools/browser_agent.py
+8 -6
@@ -183,14 +183,16 @@ class BrowserAgent(Tool):
183 # collect result
184 result = await task.result()
185 answer = result.final_result()
186 +
187 + # Ensure answer is not None
188 + if answer is None:
189 + answer = "Browser task completed but no result was returned."
190 +
191 try:
187 - if answer and isinstance(answer, str) and answer.strip():
188 - answer_data = DirtyJson.parse_string(answer)
189 - answer_text = strings.dict_to_text(answer_data) # type: ignore
190 - else:
191 - answer_text = str(answer) if answer else "No result returned"
192 + answer_data = DirtyJson.parse_string(answer)
193 + answer_text = strings.dict_to_text(answer_data) # type: ignore
194 except Exception as e:
193 - answer_text = str(answer) if answer else f"Error processing result: {str(e)}"
195 + answer_text = answer
196 self.log.update(answer=answer_text)
197 return Response(message=answer, break_loop=False)
198
requirements.txt
+2 -7
@@ -11,14 +11,9 @@ flask-basicauth==0.2.0
11 flaredantic==0.1.4
12 GitPython==3.1.43
13 inputimeout==1.0.4
14 -langchain-anthropic==0.3.3
14 +langchain-core>=0.3.0
15 langchain-community==0.3.19
16 -langchain-google-genai==2.0.8
17 -langchain-groq==0.2.2
18 -langchain-huggingface==0.1.2
19 -langchain-mistralai==0.2.4
20 -langchain-ollama==0.2.2
21 -langchain-openai==0.3.1
16 +litellm==1.39.3
17 openai-whisper==20240930
18 lxml_html_clean==0.3.1
19 markdown==3.7