515 lines
16 KiB
Python
515 lines
16 KiB
Python
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,
|
||
),
|
||
)
|