feat: MCP/OpenAPI tool calling support
This commit is contained in:
committed by
GitHub
parent
a5037aa808
commit
76c6086db3
+31
-14
@@ -2,6 +2,7 @@
|
||||
|
||||
import logging
|
||||
from collections import defaultdict
|
||||
from typing import Any
|
||||
|
||||
from telegram import Update
|
||||
from telegram.constants import ParseMode
|
||||
@@ -17,17 +18,18 @@ from steward.config import Settings
|
||||
from steward.llm.client import LLMClient
|
||||
from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary
|
||||
from steward.proposals.generator import Proposal, ProposalGenerator
|
||||
from steward.tools.client import ToolClient
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Per-user conversation history for non-threaded messages (capped at _MAX_HISTORY turns).
|
||||
_history: dict[int, list[dict[str, str]]] = defaultdict(list)
|
||||
_history: dict[int, list[dict[str, Any]]] = defaultdict(list)
|
||||
_MAX_HISTORY = 20
|
||||
|
||||
# Per-thread conversation history: (chat_id, thread_id) → full history (unbounded).
|
||||
# Messages belonging to a Telegram message thread are kept in their entirety here
|
||||
# until explicitly flushed by the /flush command.
|
||||
_thread_history: dict[tuple[int, int], list[dict[str, str]]] = defaultdict(list)
|
||||
_thread_history: dict[tuple[int, int], list[dict[str, Any]]] = defaultdict(list)
|
||||
|
||||
_FLUSH_SYSTEM_PROMPT = (
|
||||
"You are Steward. The following is a complete Telegram message thread conversation. "
|
||||
@@ -318,9 +320,9 @@ async def analyse_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
|
||||
|
||||
|
||||
def _with_kb_context(
|
||||
history: list[dict[str, str]],
|
||||
history: list[dict[str, Any]],
|
||||
relevant: list[ThreadSummary],
|
||||
) -> list[dict[str, str]]:
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Prepend relevant knowledge-base summaries as a transient system context message.
|
||||
|
||||
The returned list is a *new* list — the original *history* is not mutated.
|
||||
@@ -344,10 +346,15 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
|
||||
overlap with keywords in the current message. Any matches are injected as
|
||||
transient context — they are NOT stored in the rolling history, so they do
|
||||
not permanently consume the context window.
|
||||
|
||||
If a :class:`~steward.tools.client.ToolClient` is registered in
|
||||
``context.bot_data``, the LLM is invoked with tool calling support so it
|
||||
can take actions on the configured MCP/OpenAPI tool server.
|
||||
"""
|
||||
settings: Settings = context.bot_data["settings"]
|
||||
llm: LLMClient = context.bot_data["llm"]
|
||||
store: ThreadMemoryStore = context.bot_data["thread_store"]
|
||||
tool_client: ToolClient | None = context.bot_data.get("tool_client")
|
||||
user = update.effective_user
|
||||
if user is None or not _is_allowed(user.id, settings):
|
||||
return
|
||||
@@ -359,32 +366,42 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
|
||||
key = _thread_key(update)
|
||||
if key is not None:
|
||||
# Thread message: unbounded history
|
||||
history = _thread_history[key]
|
||||
history: list[dict[str, Any]] = _thread_history[key]
|
||||
call_history = _with_kb_context(history, store.search(text))
|
||||
reply = await llm.chat(text, history=call_history)
|
||||
if tool_client is not None:
|
||||
reply = await llm.chat_with_tools(text, tool_client, history=call_history)
|
||||
else:
|
||||
reply = await llm.chat(text, history=call_history) # type: ignore[arg-type]
|
||||
history.append({"role": "user", "content": text})
|
||||
history.append({"role": "assistant", "content": reply})
|
||||
else:
|
||||
# Non-thread message: capped history per user
|
||||
history = _history[user.id]
|
||||
call_history = _with_kb_context(history, store.search(text))
|
||||
reply = await llm.chat(text, history=call_history)
|
||||
history.append({"role": "user", "content": text})
|
||||
history.append({"role": "assistant", "content": reply})
|
||||
if len(history) > _MAX_HISTORY * 2:
|
||||
_history[user.id] = history[-(_MAX_HISTORY * 2) :]
|
||||
user_history: list[dict[str, Any]] = _history[user.id]
|
||||
call_history = _with_kb_context(user_history, store.search(text))
|
||||
if tool_client is not None:
|
||||
reply = await llm.chat_with_tools(text, tool_client, history=call_history)
|
||||
else:
|
||||
reply = await llm.chat(text, history=call_history) # type: ignore[arg-type]
|
||||
user_history.append({"role": "user", "content": text})
|
||||
user_history.append({"role": "assistant", "content": reply})
|
||||
if len(user_history) > _MAX_HISTORY * 2:
|
||||
_history[user.id] = user_history[-(_MAX_HISTORY * 2) :]
|
||||
|
||||
await _send_long(update, reply)
|
||||
|
||||
|
||||
def build_application(
|
||||
settings: Settings, llm: LLMClient, thread_store: ThreadMemoryStore | None = None
|
||||
settings: Settings,
|
||||
llm: LLMClient,
|
||||
thread_store: ThreadMemoryStore | None = None,
|
||||
tool_client: ToolClient | None = None,
|
||||
) -> Application: # type: ignore[type-arg]
|
||||
"""Build and return the Telegram Application."""
|
||||
app = Application.builder().token(settings.telegram_bot_token).build()
|
||||
app.bot_data["settings"] = settings
|
||||
app.bot_data["llm"] = llm
|
||||
app.bot_data["thread_store"] = thread_store or ThreadMemoryStore(settings.thread_memory_path)
|
||||
app.bot_data["tool_client"] = tool_client # None when tools are not configured
|
||||
|
||||
app.add_handler(CommandHandler("start", start_handler))
|
||||
app.add_handler(CommandHandler("help", help_handler))
|
||||
|
||||
Reference in New Issue
Block a user