feat: integrate rate limiting into model wrapper classes

feat: integrate rate limiting into model wrapper classes - Add optional ModelConfig parameter to all model factory functions - Implement internal rate limiting in LiteLLMChatWrapper and embedding wrappers - Apply rate limiting transparently in _call, _stream, _astream, and unified_call methods - Add proper token counting for input/output rate limiting - Maintain backwards compatibility with existing model creation API - Fix browser-use bypass issue where external libraries could ignore rate limits - Share rate limiter instances across models with same provider/name - Handle async/sync rate limiting with proper event loop management Fixes issue where browser-use and other external libraries bypass Agent Zero's rate limiting by obtaining models directly from agent.get_*_model() methods.

TerminallyLazy committed Jul 27, 2025 at 18:06 UTC 482b1c82a9daef3a86777c426903fe5fd3fd10f2
3 files changed +187 -14
.gitignore
+3 -1
@@ -55,4 +55,6 @@ instruments/**
55
56 # Global rule to include .gitkeep files anywhere
57 !**/.gitkeep
58 -agent_history.gif
\ No newline at end of file
58 +agent_history.gif
59 +
60 +tests/**
agent.py
+4
@@ -576,6 +576,7 @@ class Agent:
576 return models.get_chat_model(
577 self.config.chat_model.provider,
578 self.config.chat_model.name,
579 + model_config=self.config.chat_model,
580 **self.config.chat_model.build_kwargs(),
581 )
582
@@ -583,6 +584,7 @@ class Agent:
584 return models.get_chat_model(
585 self.config.utility_model.provider,
586 self.config.utility_model.name,
587 + model_config=self.config.utility_model,
588 **self.config.utility_model.build_kwargs(),
589 )
590
@@ -590,6 +592,7 @@ class Agent:
592 return models.get_browser_model(
593 self.config.browser_model.provider,
594 self.config.browser_model.name,
595 + model_config=self.config.browser_model,
596 **self.config.browser_model.build_kwargs(),
597 )
598
@@ -597,6 +600,7 @@ class Agent:
600 return models.get_embedding_model(
601 self.config.embeddings_model.provider,
602 self.config.embeddings_model.name,
603 + model_config=self.config.embeddings_model,
604 **self.config.embeddings_model.build_kwargs(),
605 )
606
models.py
+180 -13
@@ -22,6 +22,7 @@ from python.helpers.dotenv import load_dotenv
22 from python.helpers.providers import get_provider_config
23 from python.helpers.rate_limiter import RateLimiter
24 from python.helpers.tokens import approximate_tokens
25 +import asyncio
26
27 from langchain_core.language_models.chat_models import SimpleChatModel
28 from langchain_core.outputs.chat_generation import ChatGenerationChunk
@@ -113,14 +114,38 @@ class LiteLLMChatWrapper(SimpleChatModel):
114 model_name: str
115 provider: str
116 kwargs: dict = {}
117 +
118 + class Config:
119 + arbitrary_types_allowed = True
120 + extra = "allow" # Allow extra attributes
121 + validate_assignment = False # Don't validate on assignment
122
117 - def __init__(self, model: str, provider: str, **kwargs: Any):
123 + def __init__(self, model: str, provider: str, model_config: Optional[ModelConfig] = None, **kwargs: Any):
124 model_value = f"{provider}/{model}"
125 super().__init__(model_name=model_value, provider=provider, kwargs=kwargs) # type: ignore
126 + # Set rate limiting config as instance attribute after parent init
127 + self.rate_limit_config = model_config
128
129 @property
130 def _llm_type(self) -> str:
131 return "litellm-chat"
132 +
133 + async def _apply_rate_limiting(self, input_text: str, background: bool = False):
134 + """Apply rate limiting if rate_limit_config is provided"""
135 + if not self.rate_limit_config:
136 + return None
137 +
138 + limiter = get_rate_limiter(
139 + self.rate_limit_config.provider,
140 + self.rate_limit_config.name,
141 + self.rate_limit_config.limit_requests,
142 + self.rate_limit_config.limit_input,
143 + self.rate_limit_config.limit_output,
144 + )
145 + limiter.add(input=approximate_tokens(input_text))
146 + limiter.add(requests=1)
147 + await limiter.wait()
148 + return limiter
149
150 def _convert_messages(self, messages: List[BaseMessage]) -> List[dict]:
151 result = []
@@ -177,7 +202,20 @@ class LiteLLMChatWrapper(SimpleChatModel):
202 run_manager: Optional[CallbackManagerForLLMRun] = None,
203 **kwargs: Any,
204 ) -> str:
205 + import asyncio
206 +
207 msgs = self._convert_messages(messages)
208 +
209 + # Apply rate limiting if configured
210 + if self.rate_limit_config:
211 + try:
212 + # Convert messages to text for token counting
213 + messages_text = str(msgs)
214 + asyncio.run(self._apply_rate_limiting(messages_text))
215 + except Exception as e:
216 + # Don't let rate limiting break the model call
217 + print(f"Rate limiting warning: {e}")
218 +
219 resp = completion(
220 model=self.model_name, messages=msgs, stop=stop, **{**self.kwargs, **kwargs}
221 )
@@ -191,7 +229,20 @@ class LiteLLMChatWrapper(SimpleChatModel):
229 run_manager: Optional[CallbackManagerForLLMRun] = None,
230 **kwargs: Any,
231 ) -> Iterator[ChatGenerationChunk]:
232 + import asyncio
233 +
234 msgs = self._convert_messages(messages)
235 +
236 + # Apply rate limiting if configured
237 + if self.rate_limit_config:
238 + try:
239 + # Convert messages to text for token counting
240 + messages_text = str(msgs)
241 + asyncio.run(self._apply_rate_limiting(messages_text))
242 + except Exception as e:
243 + # Don't let rate limiting break the model call
244 + print(f"Rate limiting warning: {e}")
245 +
246 for chunk in completion(
247 model=self.model_name,
248 messages=msgs,
@@ -214,6 +265,13 @@ class LiteLLMChatWrapper(SimpleChatModel):
265 **kwargs: Any,
266 ) -> AsyncIterator[ChatGenerationChunk]:
267 msgs = self._convert_messages(messages)
268 +
269 + # Apply rate limiting if configured
270 + if self.rate_limit_config:
271 + # Convert messages to text for token counting
272 + messages_text = str(msgs)
273 + await self._apply_rate_limiting(messages_text)
274 +
275 response = await acompletion(
276 model=self.model_name,
277 messages=msgs,
@@ -253,6 +311,13 @@ class LiteLLMChatWrapper(SimpleChatModel):
311 # convert to litellm format
312 msgs_conv = self._convert_messages(messages)
313
314 + # Apply rate limiting if configured
315 + limiter = None
316 + if self.rate_limit_config:
317 + # Convert messages to text for token counting
318 + messages_text = str(msgs_conv)
319 + limiter = await self._apply_rate_limiting(messages_text)
320 +
321 # call model
322 _completion = await acompletion(
323 model=self.model_name,
@@ -278,6 +343,9 @@ class LiteLLMChatWrapper(SimpleChatModel):
343 parsed["reasoning_delta"],
344 approximate_tokens(parsed["reasoning_delta"]),
345 )
346 + # Add output tokens to rate limiter if configured
347 + if limiter:
348 + limiter.add(output=approximate_tokens(parsed["reasoning_delta"]))
349 # collect response delta and call callbacks
350 if parsed["response_delta"]:
351 response += parsed["response_delta"]
@@ -288,6 +356,9 @@ class LiteLLMChatWrapper(SimpleChatModel):
356 parsed["response_delta"],
357 approximate_tokens(parsed["response_delta"]),
358 )
359 + # Add output tokens to rate limiter if configured
360 + if limiter:
361 + limiter.add(output=approximate_tokens(parsed["response_delta"]))
362
363 # return complete results
364 return response, reasoning
@@ -302,6 +373,8 @@ class BrowserCompatibleChatWrapper(LiteLLMChatWrapper):
373 def __init__(self, *args, **kwargs):
374 turn_off_logging()
375 super().__init__(*args, **kwargs)
376 + # Browser-use may expect a 'model' attribute
377 + self.model = self.model_name
378
379 def _call(
380 self,
@@ -329,12 +402,55 @@ class BrowserCompatibleChatWrapper(LiteLLMChatWrapper):
402 class LiteLLMEmbeddingWrapper(Embeddings):
403 model_name: str
404 kwargs: dict = {}
405 + rate_limit_config: Optional[ModelConfig] = None
406
333 - def __init__(self, model: str, provider: str, **kwargs: Any):
407 + def __init__(self, model: str, provider: str, model_config: Optional[ModelConfig] = None, **kwargs: Any):
408 self.model_name = f"{provider}/{model}" if provider != "openai" else model
409 self.kwargs = kwargs
410 + self.rate_limit_config = model_config
411 +
412 + def _apply_rate_limiting_sync(self, input_text: str):
413 + """Apply rate limiting synchronously if model_config is provided"""
414 + if not self.rate_limit_config:
415 + return None
416 +
417 + limiter = get_rate_limiter(
418 + self.rate_limit_config.provider,
419 + self.rate_limit_config.name,
420 + self.rate_limit_config.limit_requests,
421 + self.rate_limit_config.limit_input,
422 + self.rate_limit_config.limit_output,
423 + )
424 + limiter.add(input=approximate_tokens(input_text))
425 + limiter.add(requests=1)
426 + # Note: Embeddings typically don't have streaming, so we do synchronous rate limiting
427 + import asyncio
428 + try:
429 + asyncio.run(limiter.wait())
430 + except RuntimeError:
431 + # If we're already in an event loop, create a new thread
432 + import threading
433 + import concurrent.futures
434 +
435 + def wait_for_limiter():
436 + new_loop = asyncio.new_event_loop()
437 + asyncio.set_event_loop(new_loop)
438 + try:
439 + new_loop.run_until_complete(limiter.wait())
440 + finally:
441 + new_loop.close()
442 +
443 + with concurrent.futures.ThreadPoolExecutor() as executor:
444 + executor.submit(wait_for_limiter).result()
445 + return limiter
446
447 def embed_documents(self, texts: List[str]) -> List[List[float]]:
448 + # Apply rate limiting if configured
449 + if self.rate_limit_config:
450 + # Convert texts to combined string for token counting
451 + texts_combined = " ".join(texts)
452 + self._apply_rate_limiting_sync(texts_combined)
453 +
454 resp = embedding(model=self.model_name, input=texts, **self.kwargs)
455 return [
456 item.get("embedding") if isinstance(item, dict) else item.embedding # type: ignore
@@ -342,6 +458,10 @@ class LiteLLMEmbeddingWrapper(Embeddings):
458 ]
459
460 def embed_query(self, text: str) -> List[float]:
461 + # Apply rate limiting if configured
462 + if self.rate_limit_config:
463 + self._apply_rate_limiting_sync(text)
464 +
465 resp = embedding(model=self.model_name, input=[text], **self.kwargs)
466 item = resp.data[0] # type: ignore
467 return item.get("embedding") if isinstance(item, dict) else item.embedding # type: ignore
@@ -350,7 +470,7 @@ class LiteLLMEmbeddingWrapper(Embeddings):
470 class LocalSentenceTransformerWrapper(Embeddings):
471 """Local wrapper for sentence-transformers models to avoid HuggingFace API calls"""
472
353 - def __init__(self, provider: str, model: str, **kwargs: Any):
473 + def __init__(self, provider: str, model: str, model_config: Optional[ModelConfig] = None, **kwargs: Any):
474 # Clean common user-input mistakes
475 model = model.strip().strip('"').strip("'")
476
@@ -360,12 +480,58 @@ class LocalSentenceTransformerWrapper(Embeddings):
480
481 self.model = SentenceTransformer(model, **kwargs)
482 self.model_name = model
483 + self.rate_limit_config = model_config
484 +
485 + def _apply_rate_limiting_sync(self, input_text: str):
486 + """Apply rate limiting synchronously if model_config is provided"""
487 + if not self.rate_limit_config:
488 + return None
489 +
490 + limiter = get_rate_limiter(
491 + self.rate_limit_config.provider,
492 + self.rate_limit_config.name,
493 + self.rate_limit_config.limit_requests,
494 + self.rate_limit_config.limit_input,
495 + self.rate_limit_config.limit_output,
496 + )
497 + limiter.add(input=approximate_tokens(input_text))
498 + limiter.add(requests=1)
499 + # Note: Local models still respect rate limiting for consistency
500 + import asyncio
501 + try:
502 + asyncio.run(limiter.wait())
503 + except RuntimeError:
504 + # If we're already in an event loop, create a new thread
505 + import threading
506 + import concurrent.futures
507 +
508 + def wait_for_limiter():
509 + new_loop = asyncio.new_event_loop()
510 + asyncio.set_event_loop(new_loop)
511 + try:
512 + new_loop.run_until_complete(limiter.wait())
513 + finally:
514 + new_loop.close()
515 +
516 + with concurrent.futures.ThreadPoolExecutor() as executor:
517 + executor.submit(wait_for_limiter).result()
518 + return limiter
519
520 def embed_documents(self, texts: List[str]) -> List[List[float]]:
521 + # Apply rate limiting if configured
522 + if self.rate_limit_config:
523 + # Convert texts to combined string for token counting
524 + texts_combined = " ".join(texts)
525 + self._apply_rate_limiting_sync(texts_combined)
526 +
527 embeddings = self.model.encode(texts, convert_to_tensor=False) # type: ignore
528 return embeddings.tolist() if hasattr(embeddings, "tolist") else embeddings # type: ignore
529
530 def embed_query(self, text: str) -> List[float]:
531 + # Apply rate limiting if configured
532 + if self.rate_limit_config:
533 + self._apply_rate_limiting_sync(text)
534 +
535 embedding = self.model.encode([text], convert_to_tensor=False) # type: ignore
536 result = (
537 embedding[0].tolist() if hasattr(embedding[0], "tolist") else embedding[0]
@@ -377,6 +543,7 @@ def _get_litellm_chat(
543 cls: type = LiteLLMChatWrapper,
544 model_name: str = "",
545 provider_name: str = "",
546 + model_config: Optional[ModelConfig] = None,
547 **kwargs: Any,
548 ):
549 # use api key from kwargs or env
@@ -389,10 +556,10 @@ def _get_litellm_chat(
556 provider_name, model_name, kwargs = _adjust_call_args(
557 provider_name, model_name, kwargs
558 )
392 - return cls(provider=provider_name, model=model_name, **kwargs)
559 + return cls(provider=provider_name, model=model_name, model_config=model_config, **kwargs)
560
561
395 -def _get_litellm_embedding(model_name: str, provider_name: str, **kwargs: Any):
562 +def _get_litellm_embedding(model_name: str, provider_name: str, model_config: Optional[ModelConfig] = None, **kwargs: Any):
563 # Check if this is a local sentence-transformers model
564 if provider_name == "huggingface" and model_name.startswith(
565 "sentence-transformers/"
@@ -402,7 +569,7 @@ def _get_litellm_embedding(model_name: str, provider_name: str, **kwargs: Any):
569 provider_name, model_name, kwargs
570 )
571 return LocalSentenceTransformerWrapper(
405 - provider=provider_name, model=model_name, **kwargs
572 + provider=provider_name, model=model_name, model_config=model_config, **kwargs
573 )
574
575 # use api key from kwargs or env
@@ -415,7 +582,7 @@ def _get_litellm_embedding(model_name: str, provider_name: str, **kwargs: Any):
582 provider_name, model_name, kwargs = _adjust_call_args(
583 provider_name, model_name, kwargs
584 )
418 - return LiteLLMEmbeddingWrapper(model=model_name, provider=provider_name, **kwargs)
585 + return LiteLLMEmbeddingWrapper(model=model_name, provider=provider_name, model_config=model_config, **kwargs)
586
587
588 def _parse_chunk(chunk: Any) -> ChatChunk:
@@ -478,25 +645,25 @@ def _merge_provider_defaults(
645 return provider_name, kwargs
646
647
481 -def get_chat_model(provider: str, name: str, **kwargs: Any) -> LiteLLMChatWrapper:
648 +def get_chat_model(provider: str, name: str, model_config: Optional[ModelConfig] = None, **kwargs: Any) -> LiteLLMChatWrapper:
649 orig = provider.lower()
650 provider_name, kwargs = _merge_provider_defaults("chat", orig, kwargs)
484 - return _get_litellm_chat(LiteLLMChatWrapper, name, provider_name, **kwargs)
651 + return _get_litellm_chat(LiteLLMChatWrapper, name, provider_name, model_config, **kwargs)
652
653
654 def get_browser_model(
488 - provider: str, name: str, **kwargs: Any
655 + provider: str, name: str, model_config: Optional[ModelConfig] = None, **kwargs: Any
656 ) -> BrowserCompatibleChatWrapper:
657 orig = provider.lower()
658 provider_name, kwargs = _merge_provider_defaults("chat", orig, kwargs)
659 return _get_litellm_chat(
493 - BrowserCompatibleChatWrapper, name, provider_name, **kwargs
660 + BrowserCompatibleChatWrapper, name, provider_name, model_config, **kwargs
661 )
662
663
664 def get_embedding_model(
498 - provider: str, name: str, **kwargs: Any
665 + provider: str, name: str, model_config: Optional[ModelConfig] = None, **kwargs: Any
666 ) -> LiteLLMEmbeddingWrapper | LocalSentenceTransformerWrapper:
667 orig = provider.lower()
668 provider_name, kwargs = _merge_provider_defaults("embedding", orig, kwargs)
502 - return _get_litellm_embedding(name, provider_name, **kwargs)
669 + return _get_litellm_embedding(name, provider_name, model_config, **kwargs)