207 lines
6.5 KiB
Python
207 lines
6.5 KiB
Python
"""Core async chat completion with iterative tool calling, plus adapters."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import json
|
|
import logging
|
|
from collections.abc import Awaitable, Callable
|
|
from typing import TYPE_CHECKING, Any, cast
|
|
|
|
import openai
|
|
|
|
if TYPE_CHECKING:
|
|
from openai.types.chat import ChatCompletion, ChatCompletionMessageParam
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
ToolExecutor = Callable[[str, dict[str, str]], str]
|
|
ToolCallNotifier = Callable[[str, dict[str, str]], None | Awaitable[None]]
|
|
|
|
_chat_client: openai.AsyncOpenAI | None = None
|
|
|
|
|
|
def get_chat_client() -> openai.AsyncOpenAI:
|
|
"""Return the shared async chat client, building it once."""
|
|
global _chat_client
|
|
if _chat_client is None:
|
|
from vibe_bot.config import CHAT_ENDPOINT, CHAT_ENDPOINT_KEY
|
|
|
|
_chat_client = openai.AsyncOpenAI(
|
|
base_url=CHAT_ENDPOINT, api_key=CHAT_ENDPOINT_KEY
|
|
)
|
|
return _chat_client
|
|
|
|
|
|
async def chat_complete(
|
|
messages: list[ChatCompletionMessageParam],
|
|
*,
|
|
model: str,
|
|
max_tokens: int,
|
|
seed: int | None = None,
|
|
tools: list[dict[str, object]] | None = None,
|
|
tool_executor: ToolExecutor | None = None,
|
|
tool_call_notifier: ToolCallNotifier | None = None,
|
|
max_tool_rounds: int = 5,
|
|
timeout: float = 60.0,
|
|
) -> str:
|
|
"""Send a chat completion, optionally with iterative tool calling.
|
|
|
|
Args:
|
|
messages: The conversation messages (system/user/assistant/tool).
|
|
model: The model to use for completion.
|
|
max_tokens: Maximum number of tokens to generate.
|
|
seed: Optional sampling seed.
|
|
tools: Optional list of tool definitions in OpenAI format.
|
|
tool_executor: Sync callable (tool_name, tool_args) -> result string.
|
|
tool_call_notifier: Optional sync-or-async callback invoked before each
|
|
tool call with (tool_name, tool_args).
|
|
max_tool_rounds: Maximum tool call rounds before giving up.
|
|
timeout: Per-request timeout in seconds.
|
|
|
|
Returns:
|
|
The model's final response text, stripped of whitespace ("" on failure).
|
|
|
|
"""
|
|
client = get_chat_client()
|
|
messages = list(messages)
|
|
|
|
for _round in range(max_tool_rounds):
|
|
kwargs: dict[str, Any] = {
|
|
"model": model,
|
|
"messages": messages,
|
|
"max_tokens": max_tokens,
|
|
"timeout": timeout,
|
|
}
|
|
if seed is not None:
|
|
kwargs["seed"] = seed
|
|
if tools:
|
|
kwargs["tools"] = cast("list[Any]", tools)
|
|
|
|
response = cast(
|
|
"ChatCompletion", await client.chat.completions.create(**kwargs)
|
|
)
|
|
if not response.choices:
|
|
return ""
|
|
|
|
message = response.choices[0].message
|
|
tool_calls = message.tool_calls
|
|
if tool_calls and tool_executor is not None:
|
|
assistant_msg: dict[str, object] = {
|
|
"role": "assistant",
|
|
"content": message.content or "",
|
|
}
|
|
tool_call_dicts: list[dict[str, object]] = []
|
|
for tool_call in tool_calls:
|
|
if tool_call.type != "function":
|
|
continue
|
|
tool_call_dicts.append(
|
|
{
|
|
"id": tool_call.id,
|
|
"type": "function",
|
|
"function": {
|
|
"name": tool_call.function.name,
|
|
"arguments": tool_call.function.arguments,
|
|
},
|
|
},
|
|
)
|
|
assistant_msg["tool_calls"] = tool_call_dicts
|
|
messages.append(cast("ChatCompletionMessageParam", assistant_msg))
|
|
|
|
for tool_call in tool_calls:
|
|
if tool_call.type != "function":
|
|
continue
|
|
tool_name = tool_call.function.name
|
|
tool_args = json.loads(tool_call.function.arguments)
|
|
|
|
if tool_call_notifier is not None:
|
|
result = tool_call_notifier(tool_name, tool_args)
|
|
if result is not None:
|
|
await result
|
|
|
|
tool_result = await asyncio.to_thread(
|
|
tool_executor, tool_name, tool_args
|
|
)
|
|
messages.append(
|
|
cast(
|
|
"ChatCompletionMessageParam",
|
|
{
|
|
"role": "tool",
|
|
"tool_call_id": tool_call.id,
|
|
"content": tool_result,
|
|
},
|
|
),
|
|
)
|
|
continue
|
|
|
|
content = message.content
|
|
if content:
|
|
return content.strip()
|
|
return ""
|
|
|
|
return ""
|
|
|
|
|
|
async def chat_completion_instruct(
|
|
system_prompt: str,
|
|
user_prompt: str,
|
|
*,
|
|
model: str,
|
|
max_tokens: int = 1000,
|
|
) -> str:
|
|
"""Instruction-based completion over :func:`chat_complete`."""
|
|
messages: list[ChatCompletionMessageParam] = [
|
|
{"role": "system", "content": system_prompt},
|
|
{"role": "user", "content": user_prompt},
|
|
]
|
|
return await chat_complete(messages, model=model, max_tokens=max_tokens, seed=-1)
|
|
|
|
|
|
async def chat_completion_with_history(
|
|
system_prompt: str,
|
|
prompts: list[dict[str, str]],
|
|
*,
|
|
model: str,
|
|
max_tokens: int = 1000,
|
|
) -> str:
|
|
"""Completion with conversation history over :func:`chat_complete`."""
|
|
messages: list[ChatCompletionMessageParam] = [
|
|
cast(
|
|
"ChatCompletionMessageParam",
|
|
{"role": "system", "content": system_prompt},
|
|
),
|
|
]
|
|
messages.extend(cast("list[ChatCompletionMessageParam]", prompts))
|
|
return await chat_complete(messages, model=model, max_tokens=max_tokens, seed=-1)
|
|
|
|
|
|
async def chat_completion_with_tools(
|
|
system_prompt: str,
|
|
prompts: list[dict[str, str]],
|
|
tools: list[dict[str, object]],
|
|
tool_executor: ToolExecutor,
|
|
*,
|
|
model: str,
|
|
max_tokens: int = 1000,
|
|
max_tool_rounds: int = 5,
|
|
tool_call_notifier: ToolCallNotifier | None = None,
|
|
) -> str:
|
|
"""Tool-capable completion over :func:`chat_complete`."""
|
|
messages: list[ChatCompletionMessageParam] = [
|
|
cast(
|
|
"ChatCompletionMessageParam",
|
|
{"role": "system", "content": system_prompt},
|
|
),
|
|
]
|
|
messages.extend(cast("list[ChatCompletionMessageParam]", prompts))
|
|
return await chat_complete(
|
|
messages,
|
|
model=model,
|
|
max_tokens=max_tokens,
|
|
seed=-1,
|
|
tools=tools,
|
|
tool_executor=tool_executor,
|
|
tool_call_notifier=tool_call_notifier,
|
|
max_tool_rounds=max_tool_rounds,
|
|
)
|