Skip to content

Commit 3ac831b

Browse files
authored
refactor openai chat kwargs (#933)
1 parent 329ca64 commit 3ac831b

2 files changed

Lines changed: 42 additions & 14 deletions

File tree

src/openenv/core/llm_client.py

Lines changed: 14 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -178,6 +178,18 @@ def __init__(
178178
api_key=api_key if api_key is not None else "not-needed",
179179
)
180180

181+
def _chat_completion_kwargs(
182+
self, messages: list[dict[str, Any]], **kwargs: Any
183+
) -> dict[str, Any]:
184+
create_kwargs: dict[str, Any] = {
185+
"model": self.model,
186+
"messages": messages,
187+
self._tokens_param: kwargs.get("max_tokens", self.max_tokens),
188+
}
189+
if not self._omit_temperature:
190+
create_kwargs["temperature"] = kwargs.get("temperature", self.temperature)
191+
return create_kwargs
192+
181193
async def complete(self, prompt: str, **kwargs) -> str:
182194
"""Send a chat completion request.
183195
@@ -195,13 +207,7 @@ async def complete(self, prompt: str, **kwargs) -> str:
195207
messages.append({"role": "system", "content": self.system_prompt})
196208
messages.append({"role": "user", "content": prompt})
197209

198-
call_kwargs: dict[str, Any] = {
199-
"model": self.model,
200-
"messages": messages,
201-
self._tokens_param: kwargs.get("max_tokens", self.max_tokens),
202-
}
203-
if not self._omit_temperature:
204-
call_kwargs["temperature"] = kwargs.get("temperature", self.temperature)
210+
call_kwargs = self._chat_completion_kwargs(messages, **kwargs)
205211
response = await self._client.chat.completions.create(**call_kwargs)
206212
return response.choices[0].message.content or ""
207213

@@ -211,13 +217,7 @@ async def complete_with_tools(
211217
tools: list[dict[str, Any]],
212218
**kwargs: Any,
213219
) -> LLMResponse:
214-
create_kwargs: dict[str, Any] = {
215-
"model": self.model,
216-
"messages": messages,
217-
self._tokens_param: kwargs.get("max_tokens", self.max_tokens),
218-
}
219-
if not self._omit_temperature:
220-
create_kwargs["temperature"] = kwargs.get("temperature", self.temperature)
220+
create_kwargs = self._chat_completion_kwargs(messages, **kwargs)
221221
openai_tools = _mcp_tools_to_openai(tools)
222222
if openai_tools:
223223
create_kwargs["tools"] = openai_tools

tests/core/test_llm_client.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -236,6 +236,34 @@ async def test_no_tool_calls(self):
236236
assert result.content == "Hello there"
237237
assert result.tool_calls == []
238238

239+
@pytest.mark.asyncio
240+
async def test_kwargs_override(self):
241+
"""Keyword arguments override default temperature and max_tokens."""
242+
mock_openai = MagicMock()
243+
mock_msg = MagicMock()
244+
mock_msg.content = "Hello there"
245+
mock_msg.tool_calls = None
246+
mock_response = MagicMock()
247+
mock_response.choices = [MagicMock()]
248+
mock_response.choices[0].message = mock_msg
249+
mock_openai.chat.completions.create = AsyncMock(return_value=mock_response)
250+
251+
with patch("openenv.core.llm_client.AsyncOpenAI", return_value=mock_openai):
252+
client = OpenAIClient("http://localhost", 8000, model="gpt-4")
253+
await client.complete_with_tools(
254+
[{"role": "user", "content": "hi"}],
255+
[],
256+
temperature=0.9,
257+
max_tokens=100,
258+
)
259+
260+
mock_openai.chat.completions.create.assert_called_once_with(
261+
model="gpt-4",
262+
messages=[{"role": "user", "content": "hi"}],
263+
temperature=0.9,
264+
max_tokens=100,
265+
)
266+
239267
@pytest.mark.asyncio
240268
async def test_with_tool_calls(self):
241269
"""Response with tool calls are parsed into ToolCall objects."""

0 commit comments

Comments
 (0)