45
from pydantic import ConfigDict
46
47
48
+DEFAULT_LITELLM_GLOBAL_KWARGS: dict[str, Any] = {
49
+ "drop_params": True,
50
+}
51
+
52
+
53
+def _normalize_litellm_kwargs(values: dict[str, Any]) -> dict[str, Any]:
54
+ # Normalize .env/UI-style scalar strings into native types for LiteLLM.
55
+ result: dict[str, Any] = {}
56
+ for k, v in values.items():
57
+ if isinstance(v, str):
58
+ stripped = v.strip()
59
+ lowered = stripped.lower()
60
+ if lowered == "true":
61
+ result[k] = True
62
+ elif lowered == "false":
63
+ result[k] = False
64
+ elif lowered in ("none", "null"):
65
+ result[k] = None
66
+ else:
67
+ try:
68
+ result[k] = int(stripped)
69
+ except ValueError:
70
+ try:
71
+ result[k] = float(stripped)
72
+ except ValueError:
73
+ result[k] = v
74
+ else:
75
+ result[k] = v
76
+ return result
77
+
78
+
79
+def get_litellm_global_kwargs() -> dict[str, Any]:
80
+ kwargs = _normalize_litellm_kwargs(DEFAULT_LITELLM_GLOBAL_KWARGS)
81
+ try:
82
+ configured = settings.get_settings().get("litellm_global_kwargs", {}) # type: ignore[union-attr]
83
+ except Exception:
84
+ configured = {}
85
+ if isinstance(configured, dict):
86
+ kwargs.update(_normalize_litellm_kwargs(configured))
87
+ return kwargs
88
+
89
+
90
# keep provider logging quiet in normal operation
91
def turn_off_logging():
92
os.environ["LITELLM_LOG"] = "ERROR" # only errors
97
logging.getLogger(name).setLevel(logging.ERROR)
98
99
100
+def set_litellm_params():
101
+ global_kwargs = get_litellm_global_kwargs()
102
+ for key, value in global_kwargs.items():
103
+ setattr(litellm, key, value)
104
+ return global_kwargs
105
+
106
+
107
+def configure_litellm():
108
+ turn_off_logging()
109
+ set_litellm_params()
110
+
111
+
112
+def _merge_litellm_call_kwargs(*overrides: dict[str, Any] | None) -> dict[str, Any]:
113
+ kwargs = get_litellm_global_kwargs()
114
+ for override in overrides:
115
+ if isinstance(override, dict):
116
+ kwargs.update(override)
117
+ return kwargs
118
+
119
+
120
# init
121
load_dotenv()
60
-turn_off_logging()
122
+configure_litellm()
123
+
124
125
class ModelType(Enum):
126
CHAT = "Chat"
453
) -> str:
454
import asyncio
455
456
+ configure_litellm()
457
msgs = self._convert_messages(messages)
458
459
# Apply rate limiting if configured
460
apply_rate_limiter_sync(self.a0_model_conf, str(msgs))
461
462
# Call the model
399
- call_kwargs = _without_stream_kwarg({**self.kwargs, **kwargs})
463
+ call_kwargs = _without_stream_kwarg(
464
+ _merge_litellm_call_kwargs(self.kwargs, kwargs)
465
+ )
466
resp = completion(
467
model=self.model_name, messages=msgs, stop=stop, **call_kwargs
468
)
481
) -> Iterator[ChatGenerationChunk]:
482
import asyncio
483
484
+ configure_litellm()
485
msgs = self._convert_messages(messages)
486
487
# Apply rate limiting if configured
488
apply_rate_limiter_sync(self.a0_model_conf, str(msgs))
489
490
result = ChatGenerationResult()
424
- call_kwargs = _without_stream_kwarg({**self.kwargs, **kwargs})
491
+ call_kwargs = _without_stream_kwarg(
492
+ _merge_litellm_call_kwargs(self.kwargs, kwargs)
493
+ )
494
495
for chunk in completion(
496
model=self.model_name,
516
run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
517
**kwargs: Any,
518
) -> AsyncIterator[ChatGenerationChunk]:
519
+ configure_litellm()
520
msgs = self._convert_messages(messages)
521
522
# Apply rate limiting if configured
523
await apply_rate_limiter(self.a0_model_conf, str(msgs))
524
525
result = ChatGenerationResult()
456
- call_kwargs = _without_stream_kwarg({**self.kwargs, **kwargs})
526
+ call_kwargs = _without_stream_kwarg(
527
+ _merge_litellm_call_kwargs(self.kwargs, kwargs)
528
+ )
529
530
response = await acompletion(
531
model=self.model_name,
560
**kwargs: Any,
561
) -> Tuple[str, str]:
562
491
- turn_off_logging()
563
+ configure_litellm()
564
565
if not messages:
566
messages = []
579
)
580
581
# Prepare call kwargs and retry config (strip A0-only params before calling LiteLLM)
510
- call_kwargs: dict[str, Any] = _without_stream_kwarg({**self.kwargs, **kwargs})
582
+ call_kwargs: dict[str, Any] = _without_stream_kwarg(
583
+ _merge_litellm_call_kwargs(self.kwargs, kwargs)
584
+ )
585
max_retries: int = int(call_kwargs.pop("a0_retry_attempts", 2))
586
retry_delay_s: float = float(call_kwargs.pop("a0_retry_delay_seconds", 1.5))
587
stream = reasoning_callback is not None or response_callback is not None or tokens_callback is not None
684
self.a0_model_conf = model_config
685
686
def embed_documents(self, texts: List[str]) -> List[List[float]]:
687
+ configure_litellm()
688
# Apply rate limiting if configured
689
apply_rate_limiter_sync(self.a0_model_conf, " ".join(texts))
690
616
- resp = embedding(model=self.model_name, input=texts, **self.kwargs)
691
+ resp = embedding(
692
+ model=self.model_name,
693
+ input=texts,
694
+ **_merge_litellm_call_kwargs(self.kwargs),
695
+ )
696
return [
697
item.get("embedding") if isinstance(item, dict) else item.embedding # type: ignore
698
for item in resp.data # type: ignore
699
]
700
701
def embed_query(self, text: str) -> List[float]:
702
+ configure_litellm()
703
# Apply rate limiting if configured
704
apply_rate_limiter_sync(self.a0_model_conf, text)
705
626
- resp = embedding(model=self.model_name, input=[text], **self.kwargs)
706
+ resp = embedding(
707
+ model=self.model_name,
708
+ input=[text],
709
+ **_merge_litellm_call_kwargs(self.kwargs),
710
+ )
711
item = resp.data[0] # type: ignore
712
return item.get("embedding") if isinstance(item, dict) else item.embedding # type: ignore
713
865
def _merge_provider_defaults(
866
provider_type: ProviderModelType, original_provider: str, kwargs: dict
867
) -> tuple[str, dict]:
784
- # Normalize .env-style numeric strings (e.g., "timeout=30") into ints/floats for LiteLLM
785
- def _normalize_values(values: dict) -> dict:
786
- result: dict[str, Any] = {}
787
- for k, v in values.items():
788
- if isinstance(v, str):
789
- try:
790
- result[k] = int(v)
791
- except ValueError:
792
- try:
793
- result[k] = float(v)
794
- except ValueError:
795
- result[k] = v
796
- else:
797
- result[k] = v
798
- return result
799
-
868
provider_name = original_provider # default: unchanged
869
cfg = get_provider_config(provider_type, original_provider)
870
if cfg:
882
if key and key not in ("None", "NA"):
883
kwargs["api_key"] = key
884
817
- # Merge LiteLLM global kwargs (timeouts, stream_timeout, etc.)
818
- try:
819
- global_kwargs = settings.get_settings().get("litellm_global_kwargs", {}) # type: ignore[union-attr]
820
- except Exception:
821
- global_kwargs = {}
822
- if isinstance(global_kwargs, dict):
823
- for k, v in _normalize_values(global_kwargs).items():
824
- kwargs.setdefault(k, v)
885
+ # Merge LiteLLM global kwargs. Framework defaults are merged first, then
886
+ # configured global kwargs override those defaults; explicit provider/model
887
+ # kwargs still keep priority via setdefault.
888
+ for k, v in get_litellm_global_kwargs().items():
889
+ kwargs.setdefault(k, v)
890
891
return provider_name, kwargs
892