59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191 | class LiteLLMChat(ChatLLM):
"""LiteLLM-based chat model. Supports multiple endpoints via LiteLLM.
Args:
model: Any model string LiteLLM accepts (e.g. `"gpt-4o"`,
`"gemini/gemini-2.0-flash"`).
api_key: Credential for the endpoint. Falls back to the `LITELLM_API_KEY`
environment variable when not given, so a deployment can keep its
secret under its own name and pass it here.
api_base: Endpoint URL. Falls back to `LITELLM_API_BASE`.
"""
def __init__(
self,
model: str,
*,
api_key: str | None = None,
api_base: str | None = None,
) -> None:
self.model: str = model
self.api_key: str | None = (
api_key if api_key is not None else os.getenv("LITELLM_API_KEY")
)
self.api_base: str | None = (
api_base if api_base is not None else os.getenv("LITELLM_API_BASE")
)
self.default_params: dict = {}
@staticmethod
def _to_litellm_messages(messages: list[Message]) -> list[dict[str, str]]:
"""Convert Message objects into LiteLLM-compatible dicts."""
return [m.to_provider_dict() for m in messages]
def _build_completion_kwargs(
self,
messages: list[Message],
*,
tools: list | None = None,
response_format: dict | type[BaseModel] | None = None,
stream: bool = False,
) -> dict:
"""Build kwargs for LiteLLM completion/acompletion calls."""
kwargs: dict = {
"model": self.model,
"messages": self._to_litellm_messages(messages),
"api_base": self.api_base,
"api_key": self.api_key,
**self.default_params,
}
kwargs["stream"] = stream
if tools:
kwargs["tools"] = tools
if response_format is not None:
kwargs["response_format"] = response_format
return kwargs
def invoke(
self,
messages: list[Message],
*,
ctx: InvocationContext | None = None,
tools: list | None = None,
response_format: dict | type[BaseModel] | None = None,
) -> ChatLLMResponse:
kwargs = self._build_completion_kwargs(
messages,
tools=tools,
response_format=response_format,
)
with _translate_errors(self.model):
return to_chatllm_response(completion(**kwargs))
async def ainvoke(
self,
messages: list[Message],
*,
ctx: InvocationContext | None = None,
tools: list | None = None,
response_format: dict | type[BaseModel] | None = None,
) -> ChatLLMResponse:
kwargs = self._build_completion_kwargs(
messages,
tools=tools,
response_format=response_format,
)
with _translate_errors(self.model):
return to_chatllm_response(await acompletion(**kwargs))
def stream(
self,
messages: list[Message],
*,
ctx: InvocationContext | None = None,
tools: list | None = None,
response_format: dict | type[BaseModel] | None = None,
) -> AsyncIterator[ChatLLMStreamChunk]:
async def gen() -> AsyncIterator[ChatLLMStreamChunk]:
kwargs = self._build_completion_kwargs(
messages,
tools=tools,
response_format=response_format,
stream=True,
)
with _translate_errors(self.model):
stream = await acompletion(**kwargs)
async for chunk in stream:
choice_delta = chunk.choices[0].delta
delta = choice_delta.content or ""
reasoning_delta = (
getattr(choice_delta, "reasoning_content", None) or None
)
if delta or reasoning_delta:
yield ChatLLMStreamChunk(
delta=MessageChunk(role="assistant", delta=delta),
reasoning_delta=reasoning_delta,
)
for tc in getattr(choice_delta, "tool_calls", None) or []:
yield ChatLLMStreamChunk(
delta=MessageChunk(role="assistant", delta=""),
tool_calls={
"index": tc.index,
"id": tc.id,
"name": getattr(tc.function, "name", None),
"arguments": getattr(tc.function, "arguments", None)
or "",
},
)
return gen()
|