Files
kandrusyak_bot/assistant_bot/agent_tools.py
kandrusyak 7dc5a1f7f2
All checks were successful
quality / test (3.10) (push) Successful in 10m42s
quality / test (3.12) (push) Successful in 9m25s
feat: support updating notes and reminders
2026-07-27 19:53:08 +03:00

515 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any, Callable
from zoneinfo import ZoneInfo
from .models import ReminderParseResult
from .reminders import parse_reminder
from .storage import AssistantStorage
from .telegram_utils import parse_positive_int
from .time_utils import format_local_dt, to_utc_iso, utc_now
ToolResult = dict[str, Any]
ToolExecutor = Callable[[dict[str, Any], "ToolContext"], ToolResult]
@dataclass(frozen=True)
class ToolContext:
storage: AssistantStorage
user_id: int
chat_id: int
tz: ZoneInfo
@dataclass(frozen=True)
class AgentTool:
name: str
usage: str
description: str
execute: ToolExecutor
allowed_arguments: frozenset[str] = frozenset()
paragraph_before: bool = False
def prompt_line(self) -> str:
suffix = f" {self.usage}" if self.usage else " {}"
prefix = "\n" if self.paragraph_before else ""
return f"{prefix}- {self.name}{suffix} - {self.description}"
def coerce_positive_int(value: Any, default: int, maximum: int = 50) -> int:
try:
parsed = int(value)
except (TypeError, ValueError):
return default
if parsed <= 0:
return default
return min(parsed, maximum)
def coerce_required_text(arguments: dict[str, Any], key: str) -> str:
value = arguments.get(key, "")
return str(value).strip()
def parse_agent_reminder(
when: str,
text: str,
tz: ZoneInfo,
) -> ReminderParseResult | None:
parsed = parse_reminder(f"{when} {text}".strip(), tz)
if parsed:
return parsed
normalized_when = when.strip().replace("T", " ")
try:
parsed_dt = datetime.fromisoformat(normalized_when)
except ValueError:
return None
if parsed_dt.tzinfo is None:
parsed_dt = parsed_dt.replace(tzinfo=tz)
return ReminderParseResult(parsed_dt.astimezone(timezone.utc), text)
def get_current_datetime(
_arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
now = utc_now()
return {
"ok": True,
"local": now.astimezone(context.tz).strftime("%Y-%m-%d %H:%M:%S"),
"utc": to_utc_iso(now),
"timezone": context.tz.key,
}
def remember(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
text = coerce_required_text(arguments, "text")
if not text:
return {"ok": False, "error": "text is required"}
memory_id = context.storage.add_memory(context.user_id, text)
return {"ok": True, "id": memory_id, "text": text}
def list_memory(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
rows = context.storage.list_memories(
context.user_id,
coerce_positive_int(arguments.get("limit"), 20),
)
return {"ok": True, "items": rows}
def delete_memory(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
memory_id = parse_positive_int(str(arguments.get("id", "")))
if memory_id is None:
return {"ok": False, "error": "valid id is required"}
return {
"ok": context.storage.delete_memory(context.user_id, memory_id),
"id": memory_id,
}
def create_note(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
text = coerce_required_text(arguments, "text")
if not text:
return {"ok": False, "error": "text is required"}
note_id = context.storage.add_note(context.user_id, text)
return {"ok": True, "id": note_id, "text": text}
def list_notes(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
rows = context.storage.list_notes(
context.user_id,
coerce_positive_int(arguments.get("limit"), 20),
)
return {"ok": True, "items": rows}
def update_note(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
note_id = parse_positive_int(str(arguments.get("id", "")))
if note_id is None:
return {"ok": False, "error": "valid id is required"}
text = coerce_required_text(arguments, "text")
if not text:
return {"ok": False, "error": "text is required"}
return {
"ok": context.storage.update_note(context.user_id, note_id, text),
"id": note_id,
"text": text,
}
def delete_note(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
note_id = parse_positive_int(str(arguments.get("id", "")))
if note_id is None:
return {"ok": False, "error": "valid id is required"}
return {
"ok": context.storage.delete_note(context.user_id, note_id),
"id": note_id,
}
def create_reminder(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
when = coerce_required_text(arguments, "when")
text = coerce_required_text(arguments, "text")
if not when or not text:
return {"ok": False, "error": "when and text are required"}
parsed = parse_agent_reminder(when, text, context.tz)
if not parsed:
return {"ok": False, "error": "could not parse reminder time"}
if parsed.remind_at_utc <= utc_now():
return {"ok": False, "error": "reminder time must be in the future"}
reminder_id = context.storage.add_reminder(
context.user_id,
context.chat_id,
parsed.text,
parsed.remind_at_utc,
)
return {
"ok": True,
"id": reminder_id,
"text": parsed.text,
"remind_at": format_local_dt(parsed.remind_at_utc, context.tz),
}
def list_reminders(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
rows = context.storage.list_reminders(
context.user_id,
coerce_positive_int(arguments.get("limit"), 20),
)
for row in rows:
row["remind_at_local"] = format_local_dt(
row["remind_at"],
context.tz,
)
return {"ok": True, "items": rows}
def update_reminder(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
reminder_id = parse_positive_int(str(arguments.get("id", "")))
if reminder_id is None:
return {"ok": False, "error": "valid id is required"}
when = coerce_required_text(arguments, "when")
text = coerce_required_text(arguments, "text")
if not when and not text:
return {"ok": False, "error": "when or text is required"}
remind_at_utc = None
if when:
parsed = parse_agent_reminder(when, text or "напоминание", context.tz)
if not parsed:
return {"ok": False, "error": "could not parse reminder time"}
if parsed.remind_at_utc <= utc_now():
return {
"ok": False,
"error": "reminder time must be in the future",
}
remind_at_utc = parsed.remind_at_utc
result: ToolResult = {
"ok": context.storage.update_reminder(
context.user_id,
reminder_id,
text=text or None,
remind_at_utc=remind_at_utc,
),
"id": reminder_id,
}
if text:
result["text"] = text
if remind_at_utc is not None:
result["remind_at"] = format_local_dt(remind_at_utc, context.tz)
return result
def cancel_reminder(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
reminder_id = parse_positive_int(str(arguments.get("id", "")))
if reminder_id is None:
return {"ok": False, "error": "valid id is required"}
return {
"ok": context.storage.cancel_reminder(
context.user_id,
reminder_id,
),
"id": reminder_id,
}
def create_status(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
title = coerce_required_text(arguments, "title")
status = coerce_required_text(arguments, "status") or "open"
if not title:
return {"ok": False, "error": "title is required"}
item_id = context.storage.add_tracked_item(
context.user_id,
title,
status,
)
return {
"ok": True,
"id": item_id,
"title": title,
"status": status,
}
def list_statuses(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
rows = context.storage.list_tracked_items(
context.user_id,
coerce_positive_int(arguments.get("limit"), 30),
)
return {"ok": True, "items": rows}
def update_status(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
item_id = parse_positive_int(str(arguments.get("id", "")))
status = coerce_required_text(arguments, "status")
if item_id is None or not status:
return {"ok": False, "error": "valid id and status are required"}
return {
"ok": context.storage.set_tracked_status(
context.user_id,
item_id,
status,
),
"id": item_id,
"status": status,
}
def delete_status(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
item_id = parse_positive_int(str(arguments.get("id", "")))
if item_id is None:
return {"ok": False, "error": "valid id is required"}
return {
"ok": context.storage.delete_tracked_item(
context.user_id,
item_id,
),
"id": item_id,
}
def search_conversation(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
query = coerce_required_text(arguments, "query")
if not query:
return {"ok": False, "error": "query is required"}
limit = coerce_positive_int(arguments.get("limit"), 5, maximum=10)
hits = context.storage.search_conversation_messages(
context.user_id,
context.chat_id,
query,
limit,
)
discussions = [
{
"matched_message_id": hit["id"],
"score": hit["score"],
"messages": context.storage.conversation_message_window(
context.user_id,
context.chat_id,
int(hit["id"]),
),
}
for hit in hits
]
return {
"ok": True,
"query": query,
"discussions": discussions,
}
AGENT_TOOLS = (
AgentTool(
"get_current_datetime",
"",
"узнать текущее локальное и UTC-время.",
get_current_datetime,
),
AgentTool(
"remember",
'{"text":"..."}',
"сохранить важный долгосрочный факт о пользователе.",
remember,
allowed_arguments=frozenset({"text"}),
),
AgentTool(
"list_memory",
'{} (необязательно: "limit", по умолчанию 20)',
"показать сохраненную память.",
list_memory,
allowed_arguments=frozenset({"limit"}),
),
AgentTool(
"delete_memory",
'{"id":1}',
"удалить запись памяти.",
delete_memory,
allowed_arguments=frozenset({"id"}),
),
AgentTool(
"create_note",
'{"text":"..."}',
"сохранить заметку.",
create_note,
allowed_arguments=frozenset({"text"}),
),
AgentTool(
"list_notes",
'{} (необязательно: "limit", по умолчанию 20)',
"показать заметки.",
list_notes,
allowed_arguments=frozenset({"limit"}),
),
AgentTool(
"update_note",
'{"id":1,"text":"..."}',
(
"заменить текст существующей заметки. При дополнении сохрани прежний "
"текст и добавь к нему новые данные."
),
update_note,
allowed_arguments=frozenset({"id", "text"}),
),
AgentTool(
"delete_note",
'{"id":1}',
"удалить заметку.",
delete_note,
allowed_arguments=frozenset({"id"}),
),
AgentTool(
"create_reminder",
'{"when":"30m","text":"..."}',
(
"поставить напоминание. when можно указывать как '30m', "
"'через 2 часа', '18:30', '2026-07-16 18:30'. Если пользователь "
"говорит 'завтра/послезавтра/через неделю', сам рассчитай дату от "
"текущего локального времени и передай 'YYYY-MM-DD HH:MM'."
),
create_reminder,
allowed_arguments=frozenset({"when", "text"}),
),
AgentTool(
"list_reminders",
'{} (необязательно: "limit", по умолчанию 20)',
"показать активные напоминания.",
list_reminders,
allowed_arguments=frozenset({"limit"}),
),
AgentTool(
"update_reminder",
'{"id":1} (необязательно: "when" и/или "text"; хотя бы одно)',
(
"изменить время или полный текст активного напоминания без отмены. "
"when поддерживает те же форматы, что create_reminder."
),
update_reminder,
allowed_arguments=frozenset({"id", "when", "text"}),
),
AgentTool(
"cancel_reminder",
'{"id":1}',
"отменить напоминание.",
cancel_reminder,
allowed_arguments=frozenset({"id"}),
),
AgentTool(
"create_status",
'{"title":"..."} (необязательно: "status", по умолчанию "open")',
"начать отслеживать статус.",
create_status,
allowed_arguments=frozenset({"title", "status"}),
),
AgentTool(
"list_statuses",
'{} (необязательно: "limit", по умолчанию 30)',
"показать отслеживаемые статусы.",
list_statuses,
allowed_arguments=frozenset({"limit"}),
),
AgentTool(
"update_status",
'{"id":1,"status":"..."}',
"обновить статус.",
update_status,
allowed_arguments=frozenset({"id", "status"}),
),
AgentTool(
"delete_status",
'{"id":1}',
"удалить отслеживаемый объект.",
delete_status,
allowed_arguments=frozenset({"id"}),
),
AgentTool(
"search_conversation",
'{"query":"..."} (необязательно: "limit", по умолчанию 5)',
(
"найти старое обсуждение во всей сохраненной переписке по "
"содержательным ключевым словам."
),
search_conversation,
allowed_arguments=frozenset({"query", "limit"}),
paragraph_before=True,
),
)
AGENT_TOOL_BY_NAME = {tool.name: tool for tool in AGENT_TOOLS}
def render_agent_tool_catalog() -> str:
return "\n".join(tool.prompt_line() for tool in AGENT_TOOLS)
def execute_agent_tool(
name: str,
arguments: dict[str, Any],
storage: AssistantStorage,
user_id: int,
chat_id: int,
tz: ZoneInfo,
) -> ToolResult:
tool = AGENT_TOOL_BY_NAME.get(name.strip().lower())
if tool is None:
return {"ok": False, "error": f"unknown tool: {name}"}
unexpected_arguments = set(arguments) - tool.allowed_arguments
if unexpected_arguments:
unexpected = ", ".join(sorted(map(str, unexpected_arguments)))
return {
"ok": False,
"error": f"unexpected arguments: {unexpected}",
}
return tool.execute(
arguments,
ToolContext(
storage=storage,
user_id=user_id,
chat_id=chat_id,
tz=tz,
),
)