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