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