Revert "LiteLLM integration (#419)"

This reverts commit e93cb72b096323b0188273c4b56423fb95811665.

frdel committed Jun 10, 2025 at 20:55 UTC 4b3d3eb826b726d1996df91a4117d04da7802535
5 files changed +420 -245
agent.py
+8 -21
@@ -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, Optional, Dict, TypedDict, cast
8 +from typing import Any, Awaitable, Coroutine, Dict
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,7 +111,6 @@ class AgentContext:
111 "type": self.type.value,
112 }
113
114 -
114 @staticmethod
115 def log_to_all(
116 type: Log.Type,
@@ -612,14 +611,10 @@ class Agent:
611 self.config.utility_model, prompt.format(), background
612 )
613
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):
614 + async for chunk in (prompt | model).astream({}):
615 await self.handle_intervention() # wait for intervention and handle it, if paused
622 - content = chunk if isinstance(chunk, str) else str(chunk)
616 +
617 + content = models.parse_chunk(chunk)
618 limiter.add(output=tokens.approximate_tokens(content))
619 response += content
620
@@ -641,20 +636,13 @@ class Agent:
636 # rate limiter
637 limiter = await self.rate_limiter(self.config.chat_model, prompt.format())
638
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):
639 + async for chunk in (prompt | model).astream({}):
640 await self.handle_intervention() # wait for intervention and handle it, if paused
641
642 content = models.parse_chunk(chunk)
653 - content = chunk if isinstance(chunk, str) else str(chunk)
643 limiter.add(output=tokens.approximate_tokens(content))
644 response += content
645
657 -
646 if callback:
647 await callback(content, response)
648
@@ -725,8 +713,7 @@ class Agent:
713 if ":" in raw_tool_name:
714 tool_name, tool_method = raw_tool_name.split(":", 1)
715
728 - tool = None
729 - # tool = self.get_tool(name=tool_name, method=tool_method, args=tool_args, message=msg)
716 + tool = None # Initialize tool to None
717
718 # Try getting tool from MCP first
719 try:
models.py
+396 -210
@@ -1,24 +1,47 @@
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):
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):
45 ANTHROPIC = "Anthropic"
46 CHUTES = "Chutes"
47 DEEPSEEK = "DeepSeek"
@@ -38,199 +61,24 @@ class ModelProvider(Enum):
61 rate_limiters: dict[str, RateLimiter] = {}
62
63
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()}")
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()}")
68 or dotenv.get_dotenv_value(f"{service.upper()}_API_KEY")
50 - # or dotenv.get_dotenv_value(f"{service.upper()}_API_TOKEN")
51 - or "None"
69 + or dotenv.get_dotenv_value(
70 + f"{service.upper()}_API_TOKEN"
71 + ) # Added for CHUTES_API_TOKEN
72 + or "None"
73 )
74
75
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 -
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 +
82 def get_rate_limiter(
83 provider: ModelProvider, name: str, requests: int, input: int, output: int
84 ) -> RateLimiter:
@@ -244,4 +92,342 @@ def get_rate_limiter(
92 return limiter
93
94
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 + )
140
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
+3 -4
@@ -33,12 +33,11 @@ class Tool:
33 PrintStyle().print()
34
35 async def after_execution(self, response: Response, **kwargs):
36 - message = response.message or ""
37 - text = message.strip()
36 + text = response.message.strip()
37 self.agent.hist_add_tool_result(self.name, text)
38 PrintStyle(font_color="#1B4F72", background_color="white", padding=True, bold=True).print(f"{self.agent.agent_name}: Response from tool '{self.name}'")
40 - PrintStyle(font_color="#85C1E9").print(message)
41 - self.log.update(content=message)
39 + PrintStyle(font_color="#85C1E9").print(response.message)
40 + self.log.update(content=response.message)
41
42 def get_log_object(self):
43 if self.method:
python/tools/browser_agent.py
+6 -8
@@ -183,16 +183,14 @@ 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 -
186 try:
192 - answer_data = DirtyJson.parse_string(answer)
193 - answer_text = strings.dict_to_text(answer_data) # type: ignore
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 except Exception as e:
195 - answer_text = answer
193 + answer_text = str(answer) if answer else f"Error processing result: {str(e)}"
194 self.log.update(answer=answer_text)
195 return Response(message=answer, break_loop=False)
196
requirements.txt
+7 -2
@@ -11,9 +11,14 @@ flask-basicauth==0.2.0
11 flaredantic==0.1.4
12 GitPython==3.1.43
13 inputimeout==1.0.4
14 -langchain-core>=0.3.0
14 +langchain-anthropic==0.3.3
15 langchain-community==0.3.19
16 -litellm==1.39.3
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
22 openai-whisper==20240930
23 lxml_html_clean==0.3.1
24 markdown==3.7