From 84d8aaa84897dfef0ada733420718418ca7dbb24 Mon Sep 17 00:00:00 2001 From: mrleo1nid <43583385+mrleo1nid@users.noreply.github.com> Date: Thu, 17 Sep 2026 04:27:16 +0300 Subject: [PATCH] =?UTF-8?q?=D0=AD=D1=82=D0=B0=D0=BF=207:=20=D0=B4=D0=BE?= =?UTF-8?q?=D0=BB=D0=B3=D0=BE=D0=B2=D1=80=D0=B5=D0=BC=D0=B5=D0=BD=D0=BD?= =?UTF-8?q?=D0=B0=D1=8F=20=D0=BF=D0=B0=D0=BC=D1=8F=D1=82=D1=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Факты о пользователе в SQLite, в системном промпте с номерами - Инструменты remember / update_memory / forget, автоматическое запоминание отключается - Цикл вызова инструментов со стримингом (до 5 кругов), проверка и приведение аргументов - Откат без инструментов для моделей, которые их не поддерживают, с одним предупреждением - Действия в журнале, вкладка «Память» в настройках - Фейковый OpenAI-совместимый сервер для тестов, тесты полного цикла через Assistant Co-Authored-By: Claude Opus 5 --- README.md | 13 ++ src/agr_assistent/app.py | 15 +- src/agr_assistent/config.py | 16 +- src/agr_assistent/core/assistant.py | 128 ++++++++++++-- src/agr_assistent/core/context.py | 17 +- src/agr_assistent/core/memory.py | 220 ++++++++++++++++++++++++ src/agr_assistent/default_config.yaml | 6 + src/agr_assistent/llm/client.py | 85 +++++++-- src/agr_assistent/llm/tools.py | 146 ++++++++++++++++ src/agr_assistent/ui/chat_window.py | 22 ++- src/agr_assistent/ui/settings_dialog.py | 91 +++++++++- tests/__init__.py | 0 tests/conftest.py | 11 ++ tests/fake_llm.py | 90 ++++++++++ tests/test_assistant_flow.py | 141 +++++++++++++++ tests/test_llm_client.py | 67 ++++++++ tests/test_memory.py | 73 ++++++++ tests/test_tools.py | 64 +++++++ 18 files changed, 1168 insertions(+), 37 deletions(-) create mode 100644 src/agr_assistent/core/memory.py create mode 100644 src/agr_assistent/llm/tools.py create mode 100644 tests/__init__.py create mode 100644 tests/fake_llm.py create mode 100644 tests/test_assistant_flow.py create mode 100644 tests/test_llm_client.py create mode 100644 tests/test_memory.py create mode 100644 tests/test_tools.py diff --git a/README.md b/README.md index adae95d..aaf3675 100644 --- a/README.md +++ b/README.md @@ -6,6 +6,7 @@ - Каждый запрос отдельный — без бесконечной истории диалога; короткие уточнения («а завтра?») в течение пары минут после ответа видят предыдущие вопросы - Журнал запросов и ответов со стримингом +- Долговременная память: «запомни…», «забудь…», а устойчивые факты о вас модель сохраняет сама - Озвучка ответов голосом Silero: фразы проговариваются по мере генерации, блоки кода пропускаются - Голосовой ввод по глобальной горячей клавише: faster-whisper на видеокарте, конец фразы по паузе (Silero VAD) - Слово активации («ассистент») через Vosk — без нажатия клавиш @@ -37,6 +38,18 @@ llm: Голос и модель задаются в секции `tts` конфига; озвучку можно выключить в меню значка. Silero читает только кириллицу: числа переводятся в слова, латиница пропускается. +### Память + +Скажите «Запомни, что у меня Škoda Octavia» или «Забудь про машину». Если включено +автоматическое запоминание (`memory.auto_save`), модель сама сохраняет устойчивые факты: +имя, близких, технику, предпочтения. Факты хранятся локально в +`%LOCALAPPDATA%\agr-assistent\memory.sqlite3` и добавляются к каждому запросу; посмотреть, +исправить и удалить их можно в настройках на вкладке «Память». + +Память работает через вызов инструментов, поэтому модель должна их поддерживать +(например, `qwen2.5:7b` в Ollama или большинство моделей OpenRouter). С другими моделями +ассистент просто отвечает без памяти и один раз предупреждает об этом. + ### Голосовой ввод Нажмите `Win+Alt+Space` (настраивается в `voice.hotkey`), дождитесь короткого сигнала и говорите — diff --git a/src/agr_assistent/app.py b/src/agr_assistent/app.py index 10d9250..4296b4d 100644 --- a/src/agr_assistent/app.py +++ b/src/agr_assistent/app.py @@ -20,6 +20,7 @@ from agr_assistent.audio.recorder import SpeechRecorder from agr_assistent.audio.wakeword import VoskWakeWord from agr_assistent.config import AppConfig, ConfigError, data_dir, load_config from agr_assistent.core.assistant import Assistant +from agr_assistent.core.memory import MemoryStore from agr_assistent.core.settings import Settings from agr_assistent.core.speech import Speaker from agr_assistent.core.voice import VoiceInput @@ -103,7 +104,16 @@ def main(argv: list[str] | None = None) -> int: VoskWakeWord(config.wake_word, models_dir), enabled=config.wake_word.enabled ) - assistant = Assistant(config.llm, speaker, voice, wake_word) + memory = MemoryStore(data_dir() / "memory.sqlite3") + app.aboutToQuit.connect(memory.close) + assistant = Assistant( + config.llm, + speaker, + voice, + wake_word, + memory=memory, + memory_auto_save=config.memory.auto_save, + ) app.setWindowIcon(state_icon(assistant.state)) window = ChatWindow(assistant) instance.activated.connect(window.show_and_raise) @@ -113,7 +123,7 @@ def main(argv: list[str] | None = None) -> int: def open_settings() -> None: nonlocal dialog if dialog is None or not dialog.isVisible(): - dialog = SettingsDialog(settings, window) + dialog = SettingsDialog(settings, memory, window) dialog.show() dialog.raise_() dialog.activateWindow() @@ -122,6 +132,7 @@ def main(argv: list[str] | None = None) -> int: assistant.update_llm_config(new_config.llm) assistant.set_speech_enabled(new_config.tts.enabled) assistant.set_wake_word_enabled(new_config.wake_word.enabled) + assistant.set_memory_auto_save(new_config.memory.auto_save) settings.changed.connect(apply_settings) window.settings_requested.connect(open_settings) diff --git a/src/agr_assistent/config.py b/src/agr_assistent/config.py index 9eae781..3f4c5bc 100644 --- a/src/agr_assistent/config.py +++ b/src/agr_assistent/config.py @@ -85,6 +85,11 @@ class WakeWordConfig: model: str +@dataclass +class MemoryConfig: + auto_save: bool + + @dataclass class UIConfig: start_minimized: bool @@ -97,6 +102,7 @@ class AppConfig: stt: STTConfig voice: VoiceConfig wake_word: WakeWordConfig + memory: MemoryConfig ui: UIConfig path: Path @@ -254,6 +260,7 @@ def parse_config(data: dict[str, Any], path: Path) -> AppConfig: phrases=[str(phrase) for phrase in phrases if str(phrase).strip()], model=str(wake_data["model"]), ) + memory = MemoryConfig(auto_save=bool(data["memory"]["auto_save"])) ui = UIConfig(start_minimized=bool(data["ui"]["start_minimized"])) except KeyError as exc: raise ConfigError(f"{path}: отсутствует ключ {exc}") from exc @@ -268,7 +275,14 @@ def parse_config(data: dict[str, Any], path: Path) -> AppConfig: if not wake_word.phrases: raise ConfigError(f"{path}: wake_word.phrases не может быть пустым") return AppConfig( - llm=llm, tts=tts, stt=stt, voice=voice, wake_word=wake_word, ui=ui, path=path + llm=llm, + tts=tts, + stt=stt, + voice=voice, + wake_word=wake_word, + memory=memory, + ui=ui, + path=path, ) diff --git a/src/agr_assistent/core/assistant.py b/src/agr_assistent/core/assistant.py index 6f9c368..1607450 100644 --- a/src/agr_assistent/core/assistant.py +++ b/src/agr_assistent/core/assistant.py @@ -6,18 +6,31 @@ import logging import threading from datetime import datetime from enum import Enum +from typing import Any from PySide6.QtCore import QObject, Signal, Slot from agr_assistent.config import LLMConfig -from agr_assistent.core.context import FollowUpContext, Message, build_messages +from agr_assistent.core.context import FollowUpContext, build_messages +from agr_assistent.core.memory import MemoryStore, memory_prompt, memory_tools from agr_assistent.core.speech import Speaker from agr_assistent.core.voice import VoiceInput from agr_assistent.core.wake import WakeWordListener -from agr_assistent.llm.client import LLMClient, LLMError +from agr_assistent.llm.client import ( + LLMClient, + LLMError, + TextDelta, + ToolCall, + ToolCalls, + ToolsNotSupportedError, +) +from agr_assistent.llm.tools import ToolRegistry log = logging.getLogger(__name__) +# Сколько раз подряд модель может вызвать инструменты в одном ответе +MAX_TOOL_ROUNDS = 5 + class AssistantState(Enum): IDLE = "idle" @@ -39,10 +52,13 @@ class Assistant(QObject): reply_chunk = Signal(str) reply_finished = Signal(str) # полный (возможно, прерванный) текст ответа error_occurred = Signal(str) + tool_executed = Signal(str, bool) # текст для журнала; успешно ли journal_cleared = Signal() # Мост из фонового потока в главный; int — номер генерации _worker_chunk = Signal(int, str) + _worker_tool = Signal(int, str, bool) + _worker_notice = Signal(int, str) _worker_failed = Signal(int, str) _worker_done = Signal(int) @@ -52,6 +68,9 @@ class Assistant(QObject): speaker: Speaker, voice: VoiceInput | None, wake_word: WakeWordListener | None = None, + *, + memory: MemoryStore | None = None, + memory_auto_save: bool = True, parent: QObject | None = None, ) -> None: super().__init__(parent) @@ -59,6 +78,11 @@ class Assistant(QObject): self._speaker = speaker self._voice = voice self._wake_word = wake_word + self._memory = memory + self._memory_auto_save = memory_auto_save + self._tools = ToolRegistry(memory_tools(memory) if memory is not None else ()) + # Модели, которые отказались работать с инструментами: больше не предлагаем им инструменты + self._models_without_tools: set[tuple[str, str]] = set() self._client: LLMClient | None = None self._context = FollowUpContext(config.follow_up_seconds) self._state = AssistantState.IDLE @@ -69,6 +93,8 @@ class Assistant(QObject): self._reply_parts: list[str] = [] self._worker_chunk.connect(self._on_worker_chunk) + self._worker_tool.connect(self._on_worker_tool) + self._worker_notice.connect(self._on_worker_notice) self._worker_failed.connect(self._on_worker_failed) self._worker_done.connect(self._on_worker_done) speaker.playback_started.connect(self._update_state) @@ -166,6 +192,9 @@ class Assistant(QObject): log.info("Настройки LLM обновлены: %s (%s)", config.provider, self.model_name) self.provider_changed.emit(config.provider) + def set_memory_auto_save(self, enabled: bool) -> None: + self._memory_auto_save = enabled + def set_speech_enabled(self, enabled: bool) -> None: if enabled == self._speaker.enabled: return @@ -191,13 +220,20 @@ class Assistant(QObject): self._generating = True self._request = text self._reply_parts = [] - messages = build_messages(self._config.system_prompt, context, text, datetime.now()) + sections = [] + if self._memory is not None: + sections.append(memory_prompt(self._memory.facts(), self._memory_auto_save)) + messages = build_messages( + self._config.system_prompt, context, text, datetime.now(), sections + ) + model_key = (self._config.active_provider.base_url, self._config.active_provider.model) + use_tools = len(self._tools) > 0 and model_key not in self._models_without_tools self._speaker.begin() self._update_state() self.reply_started.emit() threading.Thread( target=self._run_reply, - args=(self._generation, client, messages), + args=(self._generation, client, messages, use_tools, model_key), name=f"llm-reply-{self._generation}", daemon=True, ).start() @@ -226,13 +262,40 @@ class Assistant(QObject): ) return self._client - def _run_reply(self, generation: int, client: LLMClient, messages: list[Message]) -> None: - """Выполняется в фоновом потоке.""" + def _run_reply( + self, + generation: int, + client: LLMClient, + messages: list[dict[str, Any]], + use_tools: bool, + model_key: tuple[str, str], + ) -> None: + """Выполняется в фоновом потоке: ответ модели и вызовы инструментов по кругу.""" try: - for piece in client.stream_chat(messages): - if generation != self._generation: - return - self._worker_chunk.emit(generation, piece) + for _round in range(MAX_TOOL_ROUNDS + 1): + # На последнем круге инструменты не предлагаем — модель обязана ответить текстом + tools = self._tools.schemas() if use_tools and _round < MAX_TOOL_ROUNDS else None + try: + calls, text = self._stream_round(generation, client, messages, tools) + except ToolsNotSupportedError as exc: + self._models_without_tools.add(model_key) + self._worker_notice.emit(generation, str(exc)) + use_tools = False + calls, text = self._stream_round(generation, client, messages, None) + if calls is None: + return # запрос отменён + if not calls: + break + messages.append(_assistant_tool_message(text, calls)) + for call in calls: + result = self._tools.execute(call.name, call.arguments) + log.info("Инструмент %s(%s): %s", call.name, call.arguments, result.content) + if generation != self._generation: + return + self._worker_tool.emit(generation, result.display, result.ok) + messages.append( + {"role": "tool", "tool_call_id": call.id, "content": result.content} + ) except LLMError as exc: self._worker_failed.emit(generation, str(exc)) except Exception as exc: @@ -241,6 +304,26 @@ class Assistant(QObject): else: self._worker_done.emit(generation) + def _stream_round( + self, + generation: int, + client: LLMClient, + messages: list[dict[str, Any]], + tools: list[dict[str, Any]] | None, + ) -> tuple[list[ToolCall] | None, str]: + """Один ответ модели. Возвращает (вызовы инструментов или None при отмене, текст).""" + calls: list[ToolCall] = [] + parts: list[str] = [] + for event in client.stream_chat(messages, tools): + if generation != self._generation: + return None, "" + if isinstance(event, TextDelta): + parts.append(event.text) + self._worker_chunk.emit(generation, event.text) + elif isinstance(event, ToolCalls): + calls = event.calls + return calls, "".join(parts) + @Slot() def _on_wake_word(self) -> None: if self._voice is not None and not self._voice.is_active: @@ -263,6 +346,16 @@ class Assistant(QObject): self._speaker.feed(piece) self.reply_chunk.emit(piece) + @Slot(int, str, bool) + def _on_worker_tool(self, generation: int, display: str, ok: bool) -> None: + if generation == self._generation: + self.tool_executed.emit(display, ok) + + @Slot(int, str) + def _on_worker_notice(self, generation: int, message: str) -> None: + if generation == self._generation: + self.error_occurred.emit(message) + @Slot(int, str) def _on_worker_failed(self, generation: int, message: str) -> None: if generation == self._generation: @@ -313,3 +406,18 @@ class Assistant(QObject): self._wake_word.resume() else: self._wake_word.pause() + + +def _assistant_tool_message(text: str, calls: list[ToolCall]) -> dict[str, Any]: + return { + "role": "assistant", + "content": text or None, + "tool_calls": [ + { + "id": call.id, + "type": "function", + "function": {"name": call.name, "arguments": call.arguments or "{}"}, + } + for call in calls + ], + } diff --git a/src/agr_assistent/core/context.py b/src/agr_assistent/core/context.py index dd36c23..cb66494 100644 --- a/src/agr_assistent/core/context.py +++ b/src/agr_assistent/core/context.py @@ -3,9 +3,10 @@ from __future__ import annotations import time -from collections.abc import Callable +from collections.abc import Callable, Iterable from dataclasses import dataclass from datetime import datetime +from typing import Any Message = dict[str, str] @@ -67,11 +68,15 @@ class FollowUpContext: def build_messages( - system_prompt: str, context: list[Message], request: str, now: datetime -) -> list[Message]: - system = "\n\n".join( - part for part in (system_prompt.strip(), f"Сейчас {format_datetime(now)}.") if part - ) + system_prompt: str, + context: list[Message], + request: str, + now: datetime, + extra_sections: Iterable[str] = (), +) -> list[dict[str, Any]]: + """extra_sections — дополнительные блоки системного промпта (память, команды).""" + parts = (system_prompt.strip(), f"Сейчас {format_datetime(now)}.", *extra_sections) + system = "\n\n".join(part.strip() for part in parts if part.strip()) return [{"role": "system", "content": system}, *context, {"role": "user", "content": request}] diff --git a/src/agr_assistent/core/memory.py b/src/agr_assistent/core/memory.py new file mode 100644 index 0000000..378ac45 --- /dev/null +++ b/src/agr_assistent/core/memory.py @@ -0,0 +1,220 @@ +"""Долговременная память: факты о пользователе в SQLite и инструменты для модели.""" + +from __future__ import annotations + +import sqlite3 +import threading +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from typing import Any + +from agr_assistent.llm.tools import Tool, ToolError, ToolResult + +SOURCE_USER = "user" +SOURCE_AUTO = "auto" + +# Больше фактов в системный промпт не кладём, чтобы не раздувать каждый запрос +MAX_PROMPT_FACTS = 200 + +_SCHEMA_VERSION = 1 + + +@dataclass(frozen=True) +class Fact: + id: int + text: str + source: str + created_at: datetime + updated_at: datetime + + @property + def is_auto(self) -> bool: + return self.source == SOURCE_AUTO + + +class MemoryStore: + """Потокобезопасное хранилище: модель пишет из фонового потока, настройки — из главного.""" + + def __init__(self, path: Path) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + self._lock = threading.Lock() + self._connection = sqlite3.connect(path, check_same_thread=False) + self._connection.row_factory = sqlite3.Row + with self._lock, self._connection: + if self._connection.execute("PRAGMA user_version").fetchone()[0] < _SCHEMA_VERSION: + self._connection.execute( + """ + CREATE TABLE IF NOT EXISTS facts ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + text TEXT NOT NULL, + source TEXT NOT NULL, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL + ) + """ + ) + self._connection.execute(f"PRAGMA user_version = {_SCHEMA_VERSION}") + + def close(self) -> None: + with self._lock: + self._connection.close() + + def facts(self) -> list[Fact]: + with self._lock: + rows = self._connection.execute("SELECT * FROM facts ORDER BY id").fetchall() + return [_fact(row) for row in rows] + + def add(self, text: str, source: str) -> Fact: + """Точный дубль (без учёта регистра) не создаётся — возвращается существующий факт.""" + text = _clean(text) + with self._lock, self._connection: + for row in self._connection.execute("SELECT * FROM facts"): + if row["text"].casefold() == text.casefold(): + return _fact(row) + now = _now() + cursor = self._connection.execute( + "INSERT INTO facts (text, source, created_at, updated_at) VALUES (?, ?, ?, ?)", + (text, source, now, now), + ) + row = self._connection.execute( + "SELECT * FROM facts WHERE id = ?", (cursor.lastrowid,) + ).fetchone() + return _fact(row) + + def update(self, fact_id: int, text: str) -> Fact | None: + text = _clean(text) + with self._lock, self._connection: + updated = self._connection.execute( + "UPDATE facts SET text = ?, updated_at = ? WHERE id = ?", (text, _now(), fact_id) + ).rowcount + row = self._connection.execute("SELECT * FROM facts WHERE id = ?", (fact_id,)).fetchone() + return _fact(row) if updated else None + + def delete(self, fact_id: int) -> Fact | None: + with self._lock, self._connection: + row = self._connection.execute("SELECT * FROM facts WHERE id = ?", (fact_id,)).fetchone() + if row is None: + return None + self._connection.execute("DELETE FROM facts WHERE id = ?", (fact_id,)) + return _fact(row) + + +def memory_prompt(facts: list[Fact], auto_save: bool) -> str: + lines = ["Долговременная память о пользователе (номер: факт):"] + if facts: + shown = facts[-MAX_PROMPT_FACTS:] + lines.extend(f"[{fact.id}] {fact.text}" for fact in shown) + else: + lines.append("пока пусто") + + lines.append("") + lines.append("Работа с памятью:") + lines.append( + "- Если пользователь просит что-то запомнить, вызови remember с requested_by_user=true " + "и коротко подтверди." + ) + if auto_save: + lines.append( + "- Если пользователь сообщает о себе устойчивый факт, который пригодится в будущем " + "(имя, близкие, работа, техника, привычки, предпочтения в ответах), сохрани его через " + "remember с requested_by_user=false, не спрашивая разрешения и не сообщая об этом. " + "Не сохраняй разовые просьбы, временное состояние и содержимое текущего вопроса." + ) + lines.append( + "- Если новый факт уточняет или противоречит сохранённому, используй update_memory " + "с номером старого факта, а не remember." + ) + lines.append("- Если пользователь просит забыть факт, вызови forget с его номером.") + lines.append( + "- Формулируй факты кратко, в третьем лице, например: «Машина пользователя — Škoda Octavia»." + ) + return "\n".join(lines) + + +def memory_tools(store: MemoryStore) -> list[Tool]: + def remember(arguments: dict[str, Any]) -> ToolResult: + text = _required_text(arguments) + source = SOURCE_USER if arguments.get("requested_by_user", True) else SOURCE_AUTO + fact = store.add(text, source) + prefix = "Запомнил" if source == SOURCE_USER else "Запомнил сам" + return ToolResult(True, f"Сохранено под номером {fact.id}", f"{prefix}: {fact.text}") + + def update_memory(arguments: dict[str, Any]) -> ToolResult: + fact = store.update(arguments["id"], _required_text(arguments)) + if fact is None: + raise ToolError(f"Факта с номером {arguments['id']} нет в памяти") + return ToolResult(True, f"Факт {fact.id} обновлён", f"Обновил в памяти: {fact.text}") + + def forget(arguments: dict[str, Any]) -> ToolResult: + fact = store.delete(arguments["id"]) + if fact is None: + raise ToolError(f"Факта с номером {arguments['id']} нет в памяти") + return ToolResult(True, f"Факт {fact.id} удалён", f"Забыл: {fact.text}") + + return [ + Tool( + name="remember", + description="Сохранить факт о пользователе в долговременную память.", + parameters={ + "type": "object", + "properties": { + "text": {"type": "string", "description": "Факт одной короткой фразой"}, + "requested_by_user": { + "type": "boolean", + "description": "true — пользователь сам попросил запомнить", + }, + }, + "required": ["text"], + }, + handler=remember, + ), + Tool( + name="update_memory", + description="Исправить или уточнить сохранённый факт.", + parameters={ + "type": "object", + "properties": { + "id": {"type": "integer", "description": "Номер факта"}, + "text": {"type": "string", "description": "Новая формулировка факта"}, + }, + "required": ["id", "text"], + }, + handler=update_memory, + ), + Tool( + name="forget", + description="Удалить факт из памяти.", + parameters={ + "type": "object", + "properties": {"id": {"type": "integer", "description": "Номер факта"}}, + "required": ["id"], + }, + handler=forget, + ), + ] + + +def _required_text(arguments: dict[str, Any]) -> str: + text = _clean(arguments.get("text", "")) + if not text: + raise ToolError("Текст факта пуст") + return text + + +def _clean(text: str) -> str: + return " ".join(str(text).split()) + + +def _now() -> str: + return datetime.now().isoformat(timespec="seconds") + + +def _fact(row: sqlite3.Row) -> Fact: + return Fact( + id=row["id"], + text=row["text"], + source=row["source"], + created_at=datetime.fromisoformat(row["created_at"]), + updated_at=datetime.fromisoformat(row["updated_at"]), + ) diff --git a/src/agr_assistent/default_config.yaml b/src/agr_assistent/default_config.yaml index 0b3e8b6..1b0b8d6 100644 --- a/src/agr_assistent/default_config.yaml +++ b/src/agr_assistent/default_config.yaml @@ -81,6 +81,12 @@ wake_word: # Модель Vosk, скачивается при первом включении (~45 МБ) model: vosk-model-small-ru-0.22 +memory: + # Модель сама сохраняет устойчивые факты о вас (имя, близкие, техника, предпочтения). + # Явные просьбы «запомни…» и «забудь…» работают всегда. Факты хранятся локально, + # просмотреть и отредактировать их можно в настройках на вкладке «Память» + auto_save: true + ui: # Запускаться сразу в трее, не показывая окно чата start_minimized: false diff --git a/src/agr_assistent/llm/client.py b/src/agr_assistent/llm/client.py index 766852b..b845e10 100644 --- a/src/agr_assistent/llm/client.py +++ b/src/agr_assistent/llm/client.py @@ -2,7 +2,10 @@ from __future__ import annotations +import re from collections.abc import Iterator +from dataclasses import dataclass +from typing import Any import openai from openai import OpenAI @@ -10,11 +13,40 @@ from openai import OpenAI from agr_assistent import APP_NAME from agr_assistent.config import ProviderConfig +# Ollama: «… does not support tools»; OpenRouter: «No endpoints found that support tool use» +_TOOLS_UNSUPPORTED = re.compile(r"support(s)?\s+tool", re.IGNORECASE) + class LLMError(Exception): """Ошибка обращения к модели с понятным пользователю текстом.""" +class ToolsNotSupportedError(LLMError): + """Модель не умеет вызывать инструменты.""" + + +@dataclass(frozen=True) +class TextDelta: + text: str + + +@dataclass(frozen=True) +class ToolCall: + id: str + name: str + arguments: str # JSON-строка, как её прислала модель + + +@dataclass(frozen=True) +class ToolCalls: + """Все вызовы инструментов ответа; приходят после окончания стрима.""" + + calls: list[ToolCall] + + +StreamEvent = TextDelta | ToolCalls + + class LLMClient: def __init__( self, provider: ProviderConfig, *, temperature: float, timeout_seconds: float @@ -35,31 +67,62 @@ class LLMClient: default_headers=headers, ) - def stream_chat(self, messages: list[dict[str, str]]) -> Iterator[str]: - """Отдаёт текст ответа по мере генерации.""" + def stream_chat( + self, messages: list[dict[str, Any]], tools: list[dict[str, Any]] | None = None + ) -> Iterator[StreamEvent]: + """Отдаёт текст по мере генерации, а в конце — запрошенные вызовы инструментов.""" provider = self._provider + extra: dict[str, Any] = {"tools": tools} if tools else {} try: stream = self._client.chat.completions.create( model=provider.model, messages=messages, # type: ignore[arg-type] temperature=self._temperature, stream=True, + **extra, ) + # Вызов инструмента приходит кусками: имя и id один раз, аргументы — частями + pending: dict[int, dict[str, str]] = {} with stream: for chunk in stream: - if chunk.choices and (text := chunk.choices[0].delta.content): - yield text + if not chunk.choices: + continue + delta = chunk.choices[0].delta + if delta.content: + yield TextDelta(delta.content) + for call in delta.tool_calls or []: + entry = pending.setdefault(call.index, {"id": "", "name": "", "arguments": ""}) + if call.id: + entry["id"] = call.id + if call.function and call.function.name and not entry["name"]: + entry["name"] = call.function.name + if call.function and call.function.arguments: + entry["arguments"] += call.function.arguments + if pending: + yield ToolCalls( + [ + ToolCall(entry["id"] or f"call_{index}", entry["name"], entry["arguments"]) + for index, entry in sorted(pending.items()) + ] + ) except openai.APITimeoutError as exc: raise LLMError(f"Превышено время ожидания ответа от {provider.base_url}") from exc except openai.APIConnectionError as exc: raise LLMError( f"Не удалось подключиться к {provider.base_url} — сервер запущен?" ) from exc - except openai.AuthenticationError as exc: - raise LLMError(f"Неверный API-ключ для провайдера '{provider.name}'") from exc - except openai.NotFoundError as exc: - raise LLMError( - f"Модель '{provider.model}' или адрес {provider.base_url} не найдены" - ) from exc except openai.APIStatusError as exc: - raise LLMError(f"Ошибка API ({exc.status_code}): {exc.message}") from exc + raise self._status_error(exc, bool(tools)) from exc + + def _status_error(self, exc: openai.APIStatusError, with_tools: bool) -> LLMError: + provider = self._provider + if with_tools and _TOOLS_UNSUPPORTED.search(str(exc.message)): + return ToolsNotSupportedError( + f"Модель '{provider.model}' не поддерживает инструменты: " + "память и команды недоступны, пока не выбрана другая модель" + ) + if isinstance(exc, openai.AuthenticationError): + return LLMError(f"Неверный API-ключ для провайдера '{provider.name}'") + if isinstance(exc, openai.NotFoundError): + return LLMError(f"Модель '{provider.model}' или адрес {provider.base_url} не найдены") + return LLMError(f"Ошибка API ({exc.status_code}): {exc.message}") diff --git a/src/agr_assistent/llm/tools.py b/src/agr_assistent/llm/tools.py new file mode 100644 index 0000000..db656cb --- /dev/null +++ b/src/agr_assistent/llm/tools.py @@ -0,0 +1,146 @@ +"""Инструменты, которые модель может вызывать: описание, проверка аргументов, выполнение.""" + +from __future__ import annotations + +import json +import logging +from collections.abc import Callable, Iterable +from dataclasses import dataclass +from typing import Any + +log = logging.getLogger(__name__) + + +class ToolError(Exception): + """Ожидаемая ошибка инструмента: текст уходит модели и в журнал.""" + + +@dataclass(frozen=True) +class ToolResult: + ok: bool + content: str # что увидит модель + display: str # что увидит пользователь в журнале + + +@dataclass(frozen=True) +class Tool: + name: str + description: str + parameters: dict[str, Any] # JSON Schema объекта аргументов + handler: Callable[[dict[str, Any]], ToolResult] + + def schema(self) -> dict[str, Any]: + return { + "type": "function", + "function": { + "name": self.name, + "description": self.description, + "parameters": self.parameters, + }, + } + + +class ToolRegistry: + def __init__(self, tools: Iterable[Tool] = ()) -> None: + self._tools: dict[str, Tool] = {} + for tool in tools: + self.register(tool) + + def __len__(self) -> int: + return len(self._tools) + + def register(self, tool: Tool) -> None: + self._tools[tool.name] = tool + + def schemas(self) -> list[dict[str, Any]]: + return [tool.schema() for tool in self._tools.values()] + + def execute(self, name: str, arguments_json: str) -> ToolResult: + tool = self._tools.get(name) + if tool is None: + return ToolResult(False, f"Инструмента {name} не существует", f"Неизвестный инструмент {name}") + try: + raw_arguments = json.loads(arguments_json or "{}") + except json.JSONDecodeError as exc: + return ToolResult(False, f"Аргументы не являются JSON: {exc}", f"{name}: некорректные аргументы") + if not isinstance(raw_arguments, dict): + return ToolResult(False, "Аргументы должны быть JSON-объектом", f"{name}: некорректные аргументы") + + arguments, errors = validate_arguments(tool.parameters, raw_arguments) + if errors: + message = "; ".join(errors) + return ToolResult(False, f"Некорректные аргументы: {message}", f"{name}: {message}") + try: + return tool.handler(arguments) + except ToolError as exc: + return ToolResult(False, str(exc), str(exc)) + except Exception as exc: + log.exception("Сбой инструмента %s", name) + return ToolResult(False, f"Ошибка выполнения: {exc}", f"{name}: ошибка — {exc}") + + +def validate_arguments( + schema: dict[str, Any], arguments: dict[str, Any] +) -> tuple[dict[str, Any], list[str]]: + """Проверяет аргументы по упрощённой JSON Schema. + + Небольшие локальные модели часто присылают числа и логические значения строками, + поэтому такие значения приводятся к нужному типу. + """ + properties: dict[str, Any] = schema.get("properties", {}) + errors = [ + f"не указан параметр {name}" for name in schema.get("required", []) if name not in arguments + ] + result: dict[str, Any] = {} + for name, value in arguments.items(): + spec = properties.get(name) + if spec is None: + continue # лишние параметры молча игнорируем + converted, error = _convert(value, spec.get("type")) + if error: + errors.append(f"{name}: {error}") + continue + if "enum" in spec and converted not in spec["enum"]: + allowed = ", ".join(str(option) for option in spec["enum"]) + errors.append(f"{name}: недопустимое значение {converted!r} (допустимы: {allowed})") + continue + result[name] = converted + return result, errors + + +def _convert(value: Any, expected: str | None) -> tuple[Any, str | None]: + if expected == "string": + if isinstance(value, (dict, list)): + return None, "ожидалась строка" + return str(value), None + if expected == "integer": + if isinstance(value, bool): + return None, "ожидалось целое число" + if isinstance(value, int): + return value, None + if isinstance(value, float) and value.is_integer(): + return int(value), None + if isinstance(value, str): + try: + return int(value.strip()), None + except ValueError: + pass + return None, "ожидалось целое число" + if expected == "number": + if isinstance(value, bool): + return None, "ожидалось число" + if isinstance(value, (int, float)): + return value, None + if isinstance(value, str): + try: + return float(value.strip().replace(",", ".")), None + except ValueError: + pass + return None, "ожидалось число" + if expected == "boolean": + if isinstance(value, bool): + return value, None + if isinstance(value, str) and value.strip().lower() in ("true", "false"): + return value.strip().lower() == "true", None + return None, "ожидалось true или false" + return value, None diff --git a/src/agr_assistent/ui/chat_window.py b/src/agr_assistent/ui/chat_window.py index 4e14b8d..d4de6f9 100644 --- a/src/agr_assistent/ui/chat_window.py +++ b/src/agr_assistent/ui/chat_window.py @@ -33,6 +33,8 @@ _ROLE_STYLES = { "user": ("Вы", "#1e88e5"), "assistant": ("Ассистент", "#43a047"), "error": ("Ошибка", "#e53935"), + "tool": ("Действие", "#8e24aa"), + "tool_failed": ("Действие не выполнено", "#e53935"), } _VOICE_BUTTON_TEXTS = { @@ -140,6 +142,7 @@ class ChatWindow(QWidget): assistant.reply_chunk.connect(self._on_reply_chunk) assistant.reply_finished.connect(self._on_reply_finished) assistant.error_occurred.connect(lambda message: self._append(_Entry("error", message))) + assistant.tool_executed.connect(self._on_tool_executed) assistant.journal_cleared.connect(self._on_journal_cleared) self._update_provider_label() @@ -188,16 +191,25 @@ class ChatWindow(QWidget): ) def _on_reply_chunk(self, piece: str) -> None: - if self._entries and self._entries[-1].role == "assistant": - self._entries[-1].text += piece - if not self._render_timer.isActive(): - self._render_timer.start() + # После действия текст ответа продолжается новой записью под ним + if not self._entries or self._entries[-1].role != "assistant": + self._entries.append(_Entry("assistant", "")) + self._entries[-1].text += piece + if not self._render_timer.isActive(): + self._render_timer.start() + + def _on_tool_executed(self, display: str, ok: bool) -> None: + self._drop_empty_reply() + self._append(_Entry("tool" if ok else "tool_failed", display)) def _on_reply_finished(self, text: str) -> None: + self._drop_empty_reply() + self._render_now() + + def _drop_empty_reply(self) -> None: last = self._entries[-1] if self._entries else None if last and last.role == "assistant" and not last.text.strip(): self._entries.pop() - self._render_now() def _on_journal_cleared(self) -> None: self._entries.clear() diff --git a/src/agr_assistent/ui/settings_dialog.py b/src/agr_assistent/ui/settings_dialog.py index a183136..3433227 100644 --- a/src/agr_assistent/ui/settings_dialog.py +++ b/src/agr_assistent/ui/settings_dialog.py @@ -8,8 +8,9 @@ from pathlib import Path from typing import Any from PySide6.QtCore import QObject, Qt, QUrl, Signal -from PySide6.QtGui import QDesktopServices +from PySide6.QtGui import QColor, QDesktopServices from PySide6.QtWidgets import ( + QAbstractItemView, QCheckBox, QComboBox, QDialog, @@ -17,8 +18,11 @@ from PySide6.QtWidgets import ( QDoubleSpinBox, QFormLayout, QHBoxLayout, + QInputDialog, QLabel, QLineEdit, + QListWidget, + QListWidgetItem, QMessageBox, QPlainTextEdit, QPushButton, @@ -31,10 +35,13 @@ from PySide6.QtWidgets import ( from agr_assistent import APP_NAME, system from agr_assistent.config import ConfigError, data_dir, expand_env, get_value +from agr_assistent.core.memory import SOURCE_USER, MemoryStore from agr_assistent.core.settings import Settings, needs_restart log = logging.getLogger(__name__) +_AUTO_FACT_COLOR = "#757575" + _TTS_SPEAKERS = ["xenia", "baya", "kseniya", "aidar", "eugene"] _STT_MODELS = ["large-v3-turbo", "large-v3", "medium", "small", "base", "tiny"] _STT_DEVICES = ["auto", "cuda", "cpu"] @@ -62,9 +69,12 @@ class _ModelListLoader(QObject): class SettingsDialog(QDialog): - def __init__(self, settings: Settings, parent: QWidget | None = None) -> None: + def __init__( + self, settings: Settings, memory: MemoryStore | None = None, parent: QWidget | None = None + ) -> None: super().__init__(parent) self._settings = settings + self._memory = memory self._raw = settings.raw() self._provider_edits: dict[str, dict[str, str]] = { name: {field: str(values.get(field) or "") for field in ("base_url", "api_key", "model")} @@ -79,6 +89,7 @@ class SettingsDialog(QDialog): tabs.addTab(self._build_llm_tab(), "Модель") tabs.addTab(self._build_speech_tab(), "Озвучка") tabs.addTab(self._build_voice_tab(), "Голосовой ввод") + tabs.addTab(self._build_memory_tab(), "Память") tabs.addTab(self._build_general_tab(), "Общие") buttons = QDialogButtonBox( @@ -206,6 +217,44 @@ class SettingsDialog(QDialog): form.addRow("Фразы", self._wake_phrases) return _page(form) + def _build_memory_tab(self) -> QWidget: + self._memory_auto_save = QCheckBox("Запоминать факты о вас автоматически") + self._memory_auto_save.setToolTip( + "Модель сама сохраняет устойчивые факты: имя, близких, технику, предпочтения. " + "Просьбы «запомни…» и «забудь…» работают всегда" + ) + self._memory_auto_save.setChecked(bool(get_value(self._raw, "memory.auto_save"))) + + self._facts = QListWidget() + self._facts.setWordWrap(True) + self._facts.itemChanged.connect(self._on_fact_edited) + add_button = QPushButton("Добавить…") + add_button.clicked.connect(self._add_fact) + delete_button = QPushButton("Удалить") + delete_button.clicked.connect(self._delete_facts) + self._facts.setSelectionMode(QAbstractItemView.SelectionMode.ExtendedSelection) + + buttons = QHBoxLayout() + buttons.addWidget(add_button) + buttons.addWidget(delete_button) + buttons.addStretch() + + hint = QLabel( + "Двойной щелчок — изменить факт. Изменения списка сохраняются сразу. " + "Серым показаны факты, которые модель сохранила сама." + ) + hint.setWordWrap(True) + + page = QWidget() + layout = QVBoxLayout(page) + layout.addWidget(self._memory_auto_save) + layout.addWidget(self._facts, 1) + layout.addLayout(buttons) + layout.addWidget(hint) + page.setEnabled(self._memory is not None) + self._reload_facts() + return page + def _build_general_tab(self) -> QWidget: self._start_minimized = QCheckBox("Запускаться свёрнутым в трей") self._start_minimized.setChecked(bool(get_value(self._raw, "ui.start_minimized"))) @@ -268,6 +317,43 @@ class SettingsDialog(QDialog): self._model.setEditText(current) self._models_status.setText(f"Доступно моделей: {len(models)}") + # --- память + + def _reload_facts(self) -> None: + if self._memory is None: + return + self._facts.blockSignals(True) + self._facts.clear() + for fact in self._memory.facts(): + item = QListWidgetItem(fact.text) + item.setData(Qt.ItemDataRole.UserRole, fact.id) + item.setFlags(item.flags() | Qt.ItemFlag.ItemIsEditable) + if fact.is_auto: + item.setForeground(QColor(_AUTO_FACT_COLOR)) + source = "сохранён автоматически" if fact.is_auto else "по вашей просьбе" + item.setToolTip(f"{source}, изменён {fact.updated_at:%d.%m.%Y %H:%M}") + self._facts.addItem(item) + self._facts.blockSignals(False) + + def _on_fact_edited(self, item: QListWidgetItem) -> None: + assert self._memory is not None + if text := item.text().strip(): + self._memory.update(item.data(Qt.ItemDataRole.UserRole), text) + self._reload_facts() + + def _add_fact(self) -> None: + assert self._memory is not None + text, ok = QInputDialog.getText(self, "Новый факт", "Что запомнить:") + if ok and text.strip(): + self._memory.add(text, SOURCE_USER) + self._reload_facts() + + def _delete_facts(self) -> None: + assert self._memory is not None + for item in self._facts.selectedItems(): + self._memory.delete(item.data(Qt.ItemDataRole.UserRole)) + self._reload_facts() + # --- сохранение def _collect(self) -> dict[str, Any]: @@ -290,6 +376,7 @@ class SettingsDialog(QDialog): "stt.language": self._stt_language.text().strip(), "wake_word.enabled": self._wake_enabled.isChecked(), "wake_word.phrases": phrases, + "memory.auto_save": self._memory_auto_save.isChecked(), "ui.start_minimized": self._start_minimized.isChecked(), } for name, fields in self._provider_edits.items(): diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/conftest.py b/tests/conftest.py index 1973fff..57e1f46 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -5,9 +5,20 @@ import pytest # Окна в тестах не показываем os.environ.setdefault("QT_QPA_PLATFORM", "offscreen") +from collections.abc import Iterator # noqa: E402 + from PySide6.QtWidgets import QApplication # noqa: E402 +from tests.fake_llm import FakeLLMServer # noqa: E402 + @pytest.fixture(scope="session") def qapp() -> QApplication: return QApplication.instance() or QApplication([]) + + +@pytest.fixture +def fake_llm() -> Iterator[FakeLLMServer]: + server = FakeLLMServer() + yield server + server.close() diff --git a/tests/fake_llm.py b/tests/fake_llm.py new file mode 100644 index 0000000..1f2ce4a --- /dev/null +++ b/tests/fake_llm.py @@ -0,0 +1,90 @@ +"""Фейковый OpenAI-совместимый сервер: отдаёт заранее заданные ответы стримом.""" + +from __future__ import annotations + +import json +import threading +from dataclasses import dataclass, field +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any + + +@dataclass +class Reply: + text: str = "" + # (имя, аргументы JSON); аргументы отдаются несколькими кусками, как у настоящих серверов + tool_calls: list[tuple[str, str]] = field(default_factory=list) + status: int = 200 + error: str = "" + + +class FakeLLMServer: + def __init__(self) -> None: + self.replies: list[Reply] = [] + self.requests: list[dict[str, Any]] = [] + server = self + + class Handler(BaseHTTPRequestHandler): + def log_message(self, *args: object) -> None: + pass + + def do_POST(self) -> None: + body = json.loads(self.rfile.read(int(self.headers["Content-Length"]))) + server.requests.append(body) + reply = server.replies.pop(0) if server.replies else Reply(text="") + if reply.status != 200: + payload = json.dumps({"error": {"message": reply.error}}).encode() + self.send_response(reply.status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(payload))) + self.end_headers() + self.wfile.write(payload) + return + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.end_headers() + for delta in _deltas(reply): + chunk = { + "id": "fake", + "object": "chat.completion.chunk", + "created": 0, + "model": body["model"], + "choices": [{"index": 0, "delta": delta, "finish_reason": None}], + } + self.wfile.write(f"data: {json.dumps(chunk)}\n\n".encode()) + self.wfile.write(b"data: [DONE]\n\n") + + self._server = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + threading.Thread(target=self._server.serve_forever, daemon=True).start() + + @property + def base_url(self) -> str: + return f"http://127.0.0.1:{self._server.server_port}/v1" + + def close(self) -> None: + self._server.shutdown() + self._server.server_close() + + +def _deltas(reply: Reply) -> list[dict[str, Any]]: + deltas: list[dict[str, Any]] = [] + for start in range(0, len(reply.text), 5): + deltas.append({"content": reply.text[start : start + 5]}) + for index, (name, arguments) in enumerate(reply.tool_calls): + middle = len(arguments) // 2 + deltas.append( + { + "tool_calls": [ + { + "index": index, + "id": f"call_{index}", + "type": "function", + "function": {"name": name, "arguments": arguments[:middle]}, + } + ] + } + ) + deltas.append( + {"tool_calls": [{"index": index, "function": {"arguments": arguments[middle:]}}]} + ) + return deltas diff --git a/tests/test_assistant_flow.py b/tests/test_assistant_flow.py new file mode 100644 index 0000000..68cf3f0 --- /dev/null +++ b/tests/test_assistant_flow.py @@ -0,0 +1,141 @@ +"""Полный цикл запроса через Assistant: модель, инструменты памяти, откат без инструментов.""" + +import time +from collections.abc import Callable +from pathlib import Path + +import numpy as np +import pytest +from PySide6.QtCore import QCoreApplication +from PySide6.QtWidgets import QApplication + +from agr_assistent.config import LLMConfig, ProviderConfig +from agr_assistent.core.assistant import Assistant +from agr_assistent.core.memory import MemoryStore +from agr_assistent.core.speech import Speaker +from tests.fake_llm import FakeLLMServer, Reply + + +class _SilentEngine: + sample_rate = 24000 + + def load(self) -> None: + pass + + def synthesize(self, text: str) -> np.ndarray: + return np.zeros(1, dtype=np.float32) + + +class _SilentPlayer: + def play(self, *args: object) -> None: + pass + + def finish(self) -> None: + pass + + def abort(self) -> None: + pass + + +def _wait_until(condition: Callable[[], bool], timeout: float = 5.0) -> None: + deadline = time.monotonic() + timeout + while not condition(): + assert time.monotonic() < deadline, "условие не выполнилось вовремя" + QCoreApplication.processEvents() + time.sleep(0.005) + + +@pytest.fixture +def memory(tmp_path: Path) -> MemoryStore: + store = MemoryStore(tmp_path / "memory.sqlite3") + yield store + store.close() + + +def _assistant(server: FakeLLMServer, memory: MemoryStore, *, auto_save: bool = True) -> Assistant: + config = LLMConfig( + provider="fake", + providers={"fake": ProviderConfig("fake", server.base_url, "key", "fake-model")}, + system_prompt="Будь краток.", + temperature=0.5, + follow_up_seconds=0, + timeout_seconds=10, + ) + speaker = Speaker(_SilentEngine(), _SilentPlayer(), enabled=False) + return Assistant(config, speaker, None, memory=memory, memory_auto_save=auto_save) + + +def test_remember_request_runs_tool_and_answers( + qapp: QApplication, fake_llm: FakeLLMServer, memory: MemoryStore +) -> None: + memory.add("Живёт в Казани", "user") + fake_llm.replies += [ + Reply(tool_calls=[("remember", '{"text": "Машина — Škoda Octavia", "requested_by_user": true}')]), + Reply(text="Запомнил."), + ] + assistant = _assistant(fake_llm, memory) + tools: list[tuple[str, bool]] = [] + finished: list[str] = [] + assistant.tool_executed.connect(lambda display, ok: tools.append((display, ok))) + assistant.reply_finished.connect(finished.append) + + assistant.send("Запомни, что у меня Škoda Octavia") + _wait_until(lambda: bool(finished)) + + assert [fact.text for fact in memory.facts()] == ["Живёт в Казани", "Машина — Škoda Octavia"] + assert tools == [("Запомнил: Машина — Škoda Octavia", True)] + assert finished == ["Запомнил."] + + first, second = fake_llm.requests + system = first["messages"][0]["content"] + assert "Живёт в Казани" in system and "requested_by_user=false" in system + assert {tool["function"]["name"] for tool in first["tools"]} == { + "remember", + "update_memory", + "forget", + } + # Второй запрос несёт вызов инструмента и его результат + assert second["messages"][-2]["tool_calls"][0]["function"]["name"] == "remember" + assert second["messages"][-1] == { + "role": "tool", + "tool_call_id": "call_0", + "content": "Сохранено под номером 2", + } + + +def test_model_without_tools_falls_back_once( + qapp: QApplication, fake_llm: FakeLLMServer, memory: MemoryStore +) -> None: + fake_llm.replies += [ + Reply(status=400, error="model does not support tools"), + Reply(text="Привет!"), + Reply(text="Снова привет!"), + ] + assistant = _assistant(fake_llm, memory) + errors: list[str] = [] + finished: list[str] = [] + assistant.error_occurred.connect(errors.append) + assistant.reply_finished.connect(finished.append) + + assistant.send("Привет") + _wait_until(lambda: len(finished) == 1) + assistant.send("Ещё раз привет") + _wait_until(lambda: len(finished) == 2) + + assert finished == ["Привет!", "Снова привет!"] + assert len(errors) == 1 and "не поддерживает инструменты" in errors[0] + assert ["tools" in request for request in fake_llm.requests] == [True, False, False] + + +def test_auto_save_can_be_disabled( + qapp: QApplication, fake_llm: FakeLLMServer, memory: MemoryStore +) -> None: + fake_llm.replies.append(Reply(text="Ок")) + assistant = _assistant(fake_llm, memory, auto_save=False) + finished: list[str] = [] + assistant.reply_finished.connect(finished.append) + + assistant.send("Меня зовут Лео") + _wait_until(lambda: bool(finished)) + + assert "requested_by_user=false" not in fake_llm.requests[0]["messages"][0]["content"] diff --git a/tests/test_llm_client.py b/tests/test_llm_client.py new file mode 100644 index 0000000..c500f96 --- /dev/null +++ b/tests/test_llm_client.py @@ -0,0 +1,67 @@ +import pytest + +from agr_assistent.config import ProviderConfig +from agr_assistent.llm.client import ( + LLMClient, + LLMError, + TextDelta, + ToolCall, + ToolCalls, + ToolsNotSupportedError, +) +from tests.fake_llm import FakeLLMServer, Reply + +_TOOLS = [{"type": "function", "function": {"name": "remember", "parameters": {"type": "object"}}}] + + +def _client(server: FakeLLMServer) -> LLMClient: + return LLMClient( + ProviderConfig("fake", server.base_url, "key", "fake-model"), + temperature=0.5, + timeout_seconds=10, + ) + + +def test_text_and_tool_calls_are_assembled(fake_llm: FakeLLMServer) -> None: + fake_llm.replies.append( + Reply(text="Сейчас запомню.", tool_calls=[("remember", '{"text": "Любит кофе"}')]) + ) + + events = list(_client(fake_llm).stream_chat([{"role": "user", "content": "привет"}], _TOOLS)) + + text = "".join(event.text for event in events if isinstance(event, TextDelta)) + assert text == "Сейчас запомню." + assert events[-1] == ToolCalls([ToolCall("call_0", "remember", '{"text": "Любит кофе"}')]) + assert fake_llm.requests[0]["tools"] == _TOOLS + + +def test_tools_are_not_sent_when_absent(fake_llm: FakeLLMServer) -> None: + fake_llm.replies.append(Reply(text="ок")) + + list(_client(fake_llm).stream_chat([{"role": "user", "content": "привет"}])) + + assert "tools" not in fake_llm.requests[0] + + +@pytest.mark.parametrize( + "message", + [ + "registry.ollama.ai/library/gemma:2b does not support tools", + "No endpoints found that support tool use", + ], +) +def test_tools_not_supported_is_recognized(fake_llm: FakeLLMServer, message: str) -> None: + fake_llm.replies.append(Reply(status=400, error=message)) + + with pytest.raises(ToolsNotSupportedError): + list(_client(fake_llm).stream_chat([{"role": "user", "content": "привет"}], _TOOLS)) + + +def test_other_api_errors_stay_generic(fake_llm: FakeLLMServer) -> None: + fake_llm.replies.append(Reply(status=400, error="context too long")) + + with pytest.raises(LLMError) as error: + list(_client(fake_llm).stream_chat([{"role": "user", "content": "привет"}], _TOOLS)) + + assert not isinstance(error.value, ToolsNotSupportedError) + assert "context too long" in str(error.value) diff --git a/tests/test_memory.py b/tests/test_memory.py new file mode 100644 index 0000000..d651298 --- /dev/null +++ b/tests/test_memory.py @@ -0,0 +1,73 @@ +from pathlib import Path + +import pytest + +from agr_assistent.core.memory import ( + SOURCE_AUTO, + SOURCE_USER, + MemoryStore, + memory_prompt, + memory_tools, +) +from agr_assistent.llm.tools import ToolRegistry + + +@pytest.fixture +def store(tmp_path: Path) -> MemoryStore: + memory = MemoryStore(tmp_path / "memory.sqlite3") + yield memory + memory.close() + + +def test_facts_survive_reopening(tmp_path: Path) -> None: + path = tmp_path / "memory.sqlite3" + first = MemoryStore(path) + fact = first.add(" Машина пользователя — Škoda Octavia ", SOURCE_USER) + first.close() + + second = MemoryStore(path) + assert [(f.id, f.text, f.source) for f in second.facts()] == [ + (fact.id, "Машина пользователя — Škoda Octavia", SOURCE_USER) + ] + second.close() + + +def test_duplicates_update_and_delete(store: MemoryStore) -> None: + fact = store.add("Зовут Лео", SOURCE_AUTO) + + assert store.add("зовут лео", SOURCE_USER).id == fact.id + assert store.update(fact.id, "Зовут Леонид").text == "Зовут Леонид" + assert store.update(999, "нет такого") is None + assert store.delete(fact.id).text == "Зовут Леонид" + assert store.delete(fact.id) is None + assert store.facts() == [] + + +def test_memory_tools(store: MemoryStore) -> None: + tools = ToolRegistry(memory_tools(store)) + + saved = tools.execute("remember", '{"text": "Любит кофе", "requested_by_user": false}') + assert saved.ok and saved.display == "Запомнил сам: Любит кофе" + fact_id = store.facts()[0].id + assert store.facts()[0].is_auto + + assert tools.execute("update_memory", f'{{"id": {fact_id}, "text": "Любит чай"}}').ok + assert store.facts()[0].text == "Любит чай" + + missing = tools.execute("forget", '{"id": 12345}') + assert not missing.ok and "12345" in missing.content + assert tools.execute("forget", f'{{"id": "{fact_id}"}}').ok + assert store.facts() == [] + assert not tools.execute("remember", '{"text": " "}').ok + + +def test_prompt_lists_facts_and_respects_auto_save(store: MemoryStore) -> None: + fact = store.add("Живёт в Казани", SOURCE_USER) + + with_auto = memory_prompt(store.facts(), auto_save=True) + without_auto = memory_prompt(store.facts(), auto_save=False) + + assert f"[{fact.id}] Живёт в Казани" in with_auto + assert "requested_by_user=false" in with_auto + assert "requested_by_user=false" not in without_auto + assert "пока пусто" in memory_prompt([], auto_save=True) diff --git a/tests/test_tools.py b/tests/test_tools.py new file mode 100644 index 0000000..6e33065 --- /dev/null +++ b/tests/test_tools.py @@ -0,0 +1,64 @@ +from typing import Any + +from agr_assistent.llm.tools import Tool, ToolError, ToolRegistry, ToolResult, validate_arguments + +_SCHEMA = { + "type": "object", + "properties": { + "level": {"type": "integer"}, + "ratio": {"type": "number"}, + "loud": {"type": "boolean"}, + "mode": {"type": "string", "enum": ["on", "off"]}, + }, + "required": ["level"], +} + + +def test_arguments_are_converted_from_strings() -> None: + arguments, errors = validate_arguments( + _SCHEMA, {"level": "30", "ratio": "0,5", "loud": "true", "mode": "on", "extra": 1} + ) + + assert errors == [] + assert arguments == {"level": 30, "ratio": 0.5, "loud": True, "mode": "on"} + + +def test_invalid_arguments_are_reported() -> None: + _arguments, errors = validate_arguments(_SCHEMA, {"ratio": "много", "loud": 1, "mode": "auto"}) + + assert "не указан параметр level" in errors + assert any(error.startswith("ratio:") for error in errors) + assert any(error.startswith("loud:") for error in errors) + assert any("недопустимое значение 'auto'" in error for error in errors) + + +def _registry() -> ToolRegistry: + def handler(arguments: dict[str, Any]) -> ToolResult: + if arguments["level"] > 100: + raise ToolError("слишком громко") + if arguments["level"] < 0: + raise RuntimeError("сломалось") + return ToolResult(True, f"ok {arguments['level']}", "готово") + + return ToolRegistry([Tool("volume", "Громкость", _SCHEMA, handler)]) + + +def test_registry_executes_tool_and_exposes_schema() -> None: + registry = _registry() + + assert registry.schemas()[0]["function"]["name"] == "volume" + assert registry.execute("volume", '{"level": 30}') == ToolResult(True, "ok 30", "готово") + + +def test_registry_turns_problems_into_failed_results() -> None: + registry = _registry() + + assert not registry.execute("missing", "{}").ok + assert not registry.execute("volume", "{not json").ok + assert not registry.execute("volume", "[1]").ok + assert not registry.execute("volume", "{}").ok + assert registry.execute("volume", '{"level": 101}') == ToolResult( + False, "слишком громко", "слишком громко" + ) + failed = registry.execute("volume", '{"level": -1}') + assert not failed.ok and "сломалось" in failed.content