Этап 7: долговременная память
- Факты о пользователе в SQLite, в системном промпте с номерами - Инструменты remember / update_memory / forget, автоматическое запоминание отключается - Цикл вызова инструментов со стримингом (до 5 кругов), проверка и приведение аргументов - Откат без инструментов для моделей, которые их не поддерживают, с одним предупреждением - Действия в журнале, вкладка «Память» в настройках - Фейковый OpenAI-совместимый сервер для тестов, тесты полного цикла через Assistant Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
542389c0f9
commit
84d8aaa848
@@ -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()
|
||||
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user