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
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 = []
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
)
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,
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,
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,
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"]
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
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,
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
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
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
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]
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
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/"
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
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:
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)