180 lines
6.5 KiB
Python
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
|