2026-07-25 12:52:42 +00:00

180 lines
6.5 KiB
Python

"""LLM client wrapper (OpenAI-compatible)."""
import json
import logging
from collections.abc import AsyncIterator
from typing import TYPE_CHECKING, Any
from openai import AsyncOpenAI
from steward.config import Settings
if TYPE_CHECKING:
from steward.tools.client import ToolClient
logger = logging.getLogger(__name__)
class LLMClient:
"""Thin async wrapper around the OpenAI chat-completions API."""
def __init__(self, settings: Settings) -> None:
self._settings = settings
self._client = AsyncOpenAI(
api_key=settings.openai_api_key,
base_url=settings.openai_base_url,
)
async def chat(
self,
user_message: str,
*,
history: list[dict[str, Any]] | None = None,
system_prompt: str | None = None,
) -> str:
"""Send a user message (with optional history) and return the assistant reply."""
system = system_prompt or self._settings.openai_system_prompt
messages: list[dict[str, Any]] = [{"role": "system", "content": system}]
if history:
messages.extend(history)
messages.append({"role": "user", "content": user_message})
logger.debug(
"LLM request: model=%s messages=%d", self._settings.openai_model, len(messages)
)
response = await self._client.chat.completions.create(
model=self._settings.openai_model,
messages=messages, # type: ignore[arg-type]
)
reply = response.choices[0].message.content or ""
logger.debug("LLM reply: %d chars", len(reply))
return reply
async def chat_with_tools(
self,
user_message: str,
tool_client: "ToolClient",
*,
history: list[dict[str, Any]] | None = None,
system_prompt: str | None = None,
max_tool_rounds: int = 10,
) -> str:
"""Chat with tool calling support (agentic loop).
Sends *user_message* to the LLM with the tool definitions supplied by
*tool_client*. If the model requests one or more tool calls, each call
is executed via *tool_client* and the results are fed back to the model.
The loop repeats until the model returns a plain-text response or
*max_tool_rounds* is reached.
Intermediate tool-call messages are **not** returned to the caller and
should **not** be stored in the persistent conversation history; only
the final assistant reply needs to be recorded alongside the original
user message.
"""
tools = await tool_client.get_tools()
system = system_prompt or self._settings.openai_system_prompt
messages: list[dict[str, Any]] = [{"role": "system", "content": system}]
if history:
messages.extend(history)
messages.append({"role": "user", "content": user_message})
logger.debug(
"LLM tool-call request: model=%s tools=%d messages=%d",
self._settings.openai_model,
len(tools),
len(messages),
)
last_content = ""
for round_num in range(max_tool_rounds):
kwargs: dict[str, Any] = {
"model": self._settings.openai_model,
"messages": messages,
}
if tools:
kwargs["tools"] = tools
kwargs["tool_choice"] = "auto"
response = await self._client.chat.completions.create(**kwargs) # type: ignore[arg-type]
choice = response.choices[0]
msg = choice.message
last_content = msg.content or ""
if not msg.tool_calls:
# No tool calls → final answer
logger.debug("LLM reply (round %d): %d chars", round_num, len(last_content))
return last_content
tool_names = [tc.function.name for tc in msg.tool_calls]
logger.info("Tool calls in round %d: %s", round_num + 1, tool_names)
# Append the assistant's tool-call turn
messages.append(
{
"role": "assistant",
"content": msg.content,
"tool_calls": [
{
"id": tc.id,
"type": "function",
"function": {
"name": tc.function.name,
"arguments": tc.function.arguments,
},
}
for tc in msg.tool_calls
],
}
)
# Execute each tool and append results
for tc in msg.tool_calls:
try:
try:
args = json.loads(tc.function.arguments)
except json.JSONDecodeError as exc:
raise ValueError(
f"Invalid JSON in tool arguments for '{tc.function.name}': {exc}"
) from exc
result = await tool_client.call(tc.function.name, args)
logger.info("Tool %s → %d chars", tc.function.name, len(result))
except Exception as exc:
result = f"Tool call failed: {exc}"
logger.warning("Tool %s failed: %s", tc.function.name, exc)
messages.append(
{
"role": "tool",
"tool_call_id": tc.id,
"content": result,
}
)
logger.warning("Reached max tool rounds (%d); returning last content", max_tool_rounds)
return last_content
async def stream(
self,
user_message: str,
*,
history: list[dict[str, Any]] | None = None,
system_prompt: str | None = None,
) -> AsyncIterator[str]:
"""Stream the assistant reply token by token."""
system = system_prompt or self._settings.openai_system_prompt
messages: list[dict[str, Any]] = [{"role": "system", "content": system}]
if history:
messages.extend(history)
messages.append({"role": "user", "content": user_message})
stream = await self._client.chat.completions.create(
model=self._settings.openai_model,
messages=messages, # type: ignore[arg-type]
stream=True,
)
async for chunk in stream:
delta = chunk.choices[0].delta.content
if delta:
yield delta