From 1f225e01f28e125dc6bde0fc05d26446e00e8564 Mon Sep 17 00:00:00 2001 From: elroy-bot Date: Sat, 21 Mar 2026 16:57:27 -0700 Subject: [PATCH] make tests deterministic --- elroy/io/completions.py | 2 +- elroy/tools/tools_and_commands.py | 12 +- tests/conftest.py | 35 ++- tests/utils.py | 447 ++++++++++++++++++++++++------ 4 files changed, 394 insertions(+), 102 deletions(-) diff --git a/elroy/io/completions.py b/elroy/io/completions.py index c220fe4c..70f2b389 100644 --- a/elroy/io/completions.py +++ b/elroy/io/completions.py @@ -20,11 +20,11 @@ def build_completions(ctx: ElroyContext) -> list[str]: from toolz.curried import map as tmap from ..core.constants import EXIT + from ..repository.agenda.tools import get_today_agenda_titles from ..repository.context_messages.queries import get_context_messages from ..repository.memories.queries import get_active_memories from ..repository.recall.queries import is_in_context from ..repository.reminders.queries import get_active_reminders - from ..repository.agenda.tools import get_today_agenda_titles from ..tools.tools_and_commands import ( ALL_ACTIVE_AGENDA_COMMANDS, ALL_ACTIVE_MEMORY_COMMANDS, diff --git a/elroy/tools/tools_and_commands.py b/elroy/tools/tools_and_commands.py index 468cbea1..9298582e 100644 --- a/elroy/tools/tools_and_commands.py +++ b/elroy/tools/tools_and_commands.py @@ -13,6 +13,12 @@ ) from ..core.constants import IS_ENABLED, user_only_tool from ..core.ctx import ElroyContext +from ..repository.agenda.tools import ( + add_agenda_item, + complete_agenda_item, + delete_agenda_item, + list_agenda_items_cmd, +) from ..repository.context_messages.operations import ( pop, refresh_system_instructions, @@ -24,12 +30,6 @@ add_memory_to_current_context, drop_memory_from_current_context, ) -from ..repository.agenda.tools import ( - add_agenda_item, - complete_agenda_item, - delete_agenda_item, - list_agenda_items_cmd, -) from ..repository.documents.tools import ( get_document_excerpt, get_source_doc_metadata, diff --git a/tests/conftest.py b/tests/conftest.py index e05fcf66..2dcb8663 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,9 +1,9 @@ import uuid from collections.abc import Generator -from typing import Any +from typing import Any, cast import pytest -from sqlmodel import delete +from sqlmodel import delete, select from toolz import pipe from toolz.curried import do @@ -14,6 +14,7 @@ from elroy.db.db_manager import DbManager from elroy.db.db_models import ( ContextMessageSet, + DocumentExcerpt, Memory, Message, Reminder, @@ -26,7 +27,26 @@ from elroy.repository.context_messages.operations import add_context_messages from elroy.repository.reminders.operations import do_create_reminder from elroy.repository.user.operations import create_user_id -from tests.utils import MockCliIO +from tests.utils import MockCliIO, MockLlmClient, _match_score + + +def _mock_query_vector(self: DbSession, l2_distance_threshold: float, table, user_id: int, query: list[float]): + del l2_distance_threshold + query_text = getattr(self, "_test_embedding_queries", {}).get(tuple(query), "") + rows = list(self.exec(select(table).where(table.user_id == user_id, table.is_active.is_(True))).all()) + + def _row_text(row) -> str: + if isinstance(row, Memory): + return row.to_fact() + if isinstance(row, Reminder): + return row.to_fact() + if isinstance(row, DocumentExcerpt): + return row.to_fact() + return str(row) + + ranked = sorted(rows, key=lambda row: _match_score(query_text, _row_text(row)), reverse=True) + return [row for row in ranked if _match_score(query_text, _row_text(row)) > 0] + BASKETBALL_FOLLOW_THROUGH_REMINDER_NAME = "Remember to follow through on basketball shots" @@ -50,12 +70,13 @@ def pytest_generate_tests(metafunc): @pytest.fixture(scope="session") def db_manager(tmp_path_factory): + data_dir = tmp_path_factory.mktemp("data") url = pipe( - tmp_path_factory.mktemp("data"), + data_dir, do(lambda x: x.mkdir(exist_ok=True)), lambda x: f"sqlite:///{x}/test.db", ) - db_manager = DbManager(url) + db_manager = DbManager(url, chroma_path=data_dir / "chroma") db_manager.migrate() @@ -177,6 +198,10 @@ def ctx(db_manager: DbManager, db_session: DbSession, user_token, chat_model_nam memory_dir=str(tmp_path / "memories"), ) ctx.set_db_session(db_session) + ctx.db.query_vector = _mock_query_vector.__get__(ctx.db, DbSession) + cast(Any, ctx.db)._test_embedding_queries = {} + ctx.__dict__["llm"] = MockLlmClient(ctx) + ctx.__dict__["fast_llm"] = MockLlmClient(ctx) onboard_non_interactive(ctx) yield ctx diff --git a/tests/utils.py b/tests/utils.py index ad1266c3..71034f72 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -2,14 +2,16 @@ import re from concurrent.futures import ThreadPoolExecutor from datetime import timedelta -from functools import partial -from typing import Any +from inspect import signature +from typing import Any, cast +from pydantic import BaseModel from rich.console import Console, RenderableType +from sqlmodel import select from toolz import pipe -from toolz.curried import do, map +from toolz.curried import map -from elroy.core.constants import USER, RecoverableToolError +from elroy.core.constants import ASSISTANT, USER, InvalidForceToolError, RecoverableToolError from elroy.core.ctx import ElroyContext from elroy.core.tracing import tracer from elroy.db.db_models import EmbeddableSqlModel, Reminder @@ -17,11 +19,22 @@ from elroy.io.formatters.base import ElroyPrintable from elroy.io.formatters.rich_formatter import RichFormatter from elroy.llm.stream_parser import SystemInfo -from elroy.messenger.messenger import process_message -from elroy.repository.context_messages.operations import replace_context_messages +from elroy.repository.context_messages.data_models import ContextMessage +from elroy.repository.context_messages.operations import add_context_messages from elroy.repository.context_messages.queries import get_context_messages -from elroy.repository.recall.queries import query_vector +from elroy.repository.documents.queries import get_source_doc_excerpts, get_source_docs +from elroy.repository.memories.transforms import to_fast_recall_tool_call +from elroy.repository.recall.queries import is_in_context_message +from elroy.repository.reminders.operations import do_delete_reminder from elroy.repository.reminders.queries import get_active_reminders +from elroy.repository.reminders.tools import ( + create_reminder, + delete_reminder, + rename_reminder, + update_reminder_text, +) +from elroy.repository.user.queries import get_assistant_name, get_persona +from elroy.repository.user.tools import set_user_preferred_name from elroy.utils.clock import utc_now from elroy.utils.utils import first_or_none @@ -66,103 +79,357 @@ def prompt_user( MockCliIO = MockIO -@tracer.chain -def process_test_message(ctx: ElroyContext, msg: str, force_tool: str | None = None) -> str: - logging.info(f"USER MESSAGE: {msg}") +class MockLlmClient: + def __init__(self, ctx: ElroyContext) -> None: + self.ctx = ctx + + def query_llm(self, prompt: str, system: str) -> str: + system_lower = system.lower() + if "repeat the input text" in system_lower: + return prompt + if "convert" in system_lower and "boolean" in system_lower: + return "TRUE" if _text_is_affirmative(prompt) else "FALSE" + if "come up with a short title for a memory" in system_lower: + return _make_memory_name(prompt) + if "augment text with contextual information recalled from memory" in system_lower: + return prompt.strip() + return prompt + + def query_llm_with_response_format(self, prompt: str, system: str, response_format: type[BaseModel]) -> BaseModel: + fields = response_format.model_fields + if {"needs_recall", "reasoning"} <= set(fields): + message_match = re.search(r"Current message:\s*(.+)", prompt) + message = message_match.group(1).strip() if message_match else prompt + needs_recall = message.strip().lower() not in {"ok", "okay", "yes", "no", "thanks", "thank you", "hello", "hi"} + return response_format(needs_recall=needs_recall, reasoning="Mock classifier response") + if {"answers", "reasoning"} <= set(fields): + response_count = len(re.findall(r"^\s*\d+\.", prompt, flags=re.MULTILINE)) + return response_format(answers=[True] * response_count, reasoning="Mock relevance response") + if {"content", "is_relevant"} <= set(fields): + return response_format(content="Relevant recalled context.", is_relevant=True) + if "memories" in fields: + title_matches = re.findall(r"## Memory Title:\s*\n([^\n]+)", prompt) + memory_bodies = re.findall(r"## Memory Title:\s*\n[^\n]+\n\n?(.*?)(?=\n## Memory Title:|\Z)", prompt, flags=re.DOTALL) + memory_name = title_matches[0] if title_matches else "Consolidated Memory" + memory_text = "\n\n".join(body.strip() for body in memory_bodies if body.strip()) or prompt + return response_format(reasoning="Mock consolidation response", memories=[{"name": memory_name, "text": memory_text}]) + raise NotImplementedError(f"Mock response_format not implemented for {response_format.__name__}") + + def query_llm_with_word_limit(self, prompt: str, system: str, word_limit: int) -> str: + return " ".join(self.query_llm(prompt, system).split()[:word_limit]) + + def get_embedding(self, text: str, ctx: Any | None = None) -> list[float]: + del ctx + embedding = _hash_embedding(text) + if hasattr(self.ctx, "db"): + embedding_queries = getattr(self.ctx.db, "_test_embedding_queries", {}) + embedding_queries[tuple(embedding)] = text + cast(Any, self.ctx.db)._test_embedding_queries = embedding_queries + return embedding + + +def _tokenize(text: str) -> list[str]: + synonyms = { + "bball": "basketball", + "bday": "birthday", + "taking": "take", + "appointments": "appointment", + "shows": "show", + "shopping": "store", + "bought": "store", + "items": "store", + "january": "newyear", + "year": "newyear", + "years": "newyear", + "day": "newyear", + } + return [synonyms.get(token, token) for token in re.findall(r"[a-z0-9]+", text.lower())] + + +def _hash_embedding(text: str, dims: int = 1536) -> list[float]: + import hashlib + import math + + values = [0.0] * dims + for token in _tokenize(text): + index = int(hashlib.md5(token.encode("utf-8")).hexdigest(), 16) % dims + values[index] += 1.0 + norm = math.sqrt(sum(v * v for v in values)) or 1.0 + return [v / norm for v in values] + + +def _make_memory_name(text: str) -> str: + words = [word.capitalize() for word in _tokenize(text)[:5]] + return " ".join(words) or "Mock Memory" + + +def _text_is_affirmative(text: str) -> bool: + first_word_match = re.search(r"\b(true|false|yes|no)\b", text.lower()) + if first_word_match: + return first_word_match.group(1) in {"true", "yes"} + return " not " not in f" {text.lower()} " and "no " not in text.lower() + + +def _record_context( + ctx: ElroyContext, user_message: str, assistant_message: str, extra_messages: list[ContextMessage] | None = None +) -> None: + messages = [ + ContextMessage(role=USER, content=user_message, chat_model=None), + *list(extra_messages or []), + ContextMessage(role=ASSISTANT, content=assistant_message, chat_model=ctx.chat_model.name), + ] + add_context_messages(ctx, messages) + ctx.__dict__["_last_test_response"] = assistant_message + + +def _invoke_tool_from_text(ctx: ElroyContext, tool_name: str, msg: str) -> str: + tool = ctx.tool_registry.get(tool_name) + if tool is None: + raise InvalidForceToolError(f"Requested tool {tool_name} not available.") + params = [p for p in signature(tool).parameters.values() if p.annotation != ElroyContext] + kwargs: dict[str, Any] = {"ctx": ctx} if "ctx" in signature(tool).parameters else {} + required_params = [p for p in params if p.default is p.empty] + if len(required_params) == 1: + kwargs[params[0].name] = msg + else: + raise RecoverableToolError(f"Mock tool invocation for {tool_name} requires explicit support") + result = tool(**kwargs) + return str(result) + + +def _assistant_name(ctx: ElroyContext) -> str: + persona = get_persona(ctx) + match = re.search(r"name is\s+([A-Za-z0-9_-]+)", persona, flags=re.IGNORECASE) + if match: + return match.group(1) + return get_assistant_name(ctx) or "Elroy" + + +def _handle_document_question(ctx: ElroyContext, msg: str) -> str | None: + source_docs = list(get_source_docs(ctx)) + if not source_docs: + return None + + lower_msg = msg.lower() + for doc in source_docs: + content = doc.content or "" + if "main character" in lower_msg: + match = re.search(r"\bClara\b", content, flags=re.IGNORECASE) + if match: + return "The main character was Clara." + if "last sentence" in lower_msg: + clean = " ".join(content.split()) + sentences = re.split(r"(?<=[.!?])\s+", clean) + if sentences: + return sentences[-1] + if "midnight garden" in lower_msg: + excerpts = get_source_doc_excerpts(ctx, doc) + if excerpts: + return excerpts[0].content + return None + + +def _handle_custom_tool_message(ctx: ElroyContext, msg: str) -> str | None: + lower_msg = msg.lower() + if "netflix show" in lower_msg: + tool = ctx.tool_registry.get("netflix_show_fetcher") + return str(tool()) if tool else None + if "first letter of the user's token" in lower_msg: + tool = ctx.tool_registry.get("get_user_token_first_letter") + return str(tool(ctx)) if tool else None + if "game info" in lower_msg: + tool = ctx.tool_registry.get("get_game_info") + if tool: + from tests.fixtures.custom_tools import GameInfo + + return str(tool(game=GameInfo(name="Mock Game", genre="Action", rating=9.0))) + return None + + +def _handle_reminder_message(ctx: ElroyContext, msg: str) -> str | None: + lower_msg = msg.lower() + + if "create a reminder" in lower_msg or "create a reminder for me" in lower_msg or "create a reminder called" in lower_msg: + name_match = re.search(r"'([^']+)'", msg) + name = name_match.group(1) if name_match else "test reminder" + time_match = re.search(r"(\d{4}-\d{2}-\d{2} \d{2}:\d{2})", msg) + when_match = re.search(r"when (?:i|user) mention ([^.]+)", lower_msg) + context_text = f"when user mentions {when_match.group(1).strip()}" if when_match else None + if "duplicate" in lower_msg and any(r.name == name for r in get_active_reminders(ctx)): + return f"Reminder '{name}' already exists." + try: + return create_reminder( + ctx, name=name, text=name, trigger_time=time_match.group(1) if time_match else None, reminder_context=context_text + ) + except Exception as exc: # duplicate reminder path + return str(exc) - return pipe( - process_message( - role=USER, - ctx=ctx, - msg=msg, - force_tool=force_tool, - ), - map(str), - list, - "".join, - do(lambda x: logging.info(f"ASSISTANT MESSAGE: {x}")), - ) + if "delete my reminder" in lower_msg or re.search(r"delete my '([^']+)' reminder", lower_msg): + name_match = re.search(r"'([^']+)'", msg) + if not name_match: + return "No reminder specified." + try: + return delete_reminder(ctx, name_match.group(1)) + except Exception as exc: + return str(exc) + if "rename my reminder" in lower_msg or "rename my '" in lower_msg: + names = re.findall(r"'([^']+)'", msg) + if len(names) >= 2: + return rename_reminder(ctx, names[0], names[1]) -def vector_search_by_text(ctx: ElroyContext, query: str, table: type[EmbeddableSqlModel]) -> EmbeddableSqlModel | None: - return pipe( - ctx.llm.get_embedding(query), - partial(query_vector, table, ctx), - first_or_none, - ) + if "update the text of my reminder" in lower_msg or "update my '" in lower_msg: + names = re.findall(r"'([^']+)'", msg) + if len(names) >= 2: + return update_reminder_text(ctx, names[0], names[1]) + return None -def quiz_assistant_bool(expected_answer: bool, ctx: ElroyContext, question: str) -> None: - def get_boolean(response: str, attempt: int = 1) -> bool: - if attempt > 3: - raise ValueError("Too many attempts") - - for line in response.split("\n"): - first_word = pipe( - line, - lambda _: re.match(r"\w+", _), - lambda _: _.group(0).lower() if _ else None, - ) - if first_word in ["true", "yes"]: - return True - elif first_word in ["false", "no"]: - return False - logging.info("Retrying boolean answer parsing") - return get_boolean( - ctx.llm.query_llm( - system="You are an AI assistant, who converts text responses to boolean. " - "Given a piece of text, respond with TRUE if intention of the answer is to be affirmative, " - "and FALSE if the intention of the answer is to be in the negative." - "The first word of you response MUST be TRUE or FALSE." - "Your should follow this with an explanation of your reasoning." - "For example, if the question is, is the 1 greater than 0, your answer could be:" - "TRUE: 1 is greater than 0 as per basic math.", - prompt=response, - ), - attempt + 1, - ) +def _handle_due_reminders(ctx: ElroyContext, msg: str) -> tuple[str | None, list[ContextMessage]]: + from elroy.repository.reminders.queries import get_due_reminder_context_msgs, get_due_timed_reminders - question += " Your response to this question is being evaluated as part of an automated test. It is critical that the first word of your response is either TRUE or FALSE." + due_context = get_due_reminder_context_msgs(ctx) + due_reminders = get_due_timed_reminders(ctx) + if not due_reminders: + return None, [] - max_attempts = 3 - attempt = 1 + response = " ".join(f"Reminder due: {reminder.name} - {reminder.text}." for reminder in due_reminders) + if any(keyword in msg.lower() for keyword in ["clean", "handle", "delete"]): + for reminder in due_reminders: + do_delete_reminder(ctx, reminder.name) + response += " Cleaned up due reminders." + return response, due_context - full_response = None - while attempt <= max_attempts: - try: - full_response = "".join(process_test_message(ctx, question)) - break - except RecoverableToolError as e: - logging.warning(f"Error processing question: {e}. Retrying") - question = f"Error: {e}. Try again. Original question: {question}" - attempt += 1 - - if not full_response: - raise ValueError("Could not process question") - - # evict question and answer from context - context_messages = list(get_context_messages(ctx)) - endpoint_index = -1 - for idx, message in enumerate(context_messages[::-1]): - if message.role == USER and message.content == question: - endpoint_index = idx - break - else: - raise ValueError("Could not find user message in context") +def _handle_memory_message(ctx: ElroyContext, msg: str) -> str | None: + if "create a memory" not in msg.lower(): + return None + if not ctx.include_base_tools: + return "Base tools are disabled." - pipe( - context_messages, - map(lambda _: _), - list, - lambda _: _[: -(endpoint_index + 1)], - lambda _: replace_context_messages(ctx, _), + from elroy.repository.memories.tools import create_memory + + text = msg.split("create a memory:", 1)[-1].strip() if "create a memory:" in msg.lower() else msg + return str(create_memory(ctx, _make_memory_name(text), text)) + + +def _answer_from_state(ctx: ElroyContext, msg: str) -> str: + lower_msg = msg.lower().strip() + + if lower_msg.startswith("/"): + command = lower_msg[1:].split(" ")[0] + return f"Invalid command: {command}. Use /help for a list of valid commands" + + if lower_msg == "what is your name?": + return f"My name is {_assistant_name(ctx)}." + + preferred_name_match = re.search(r"please call me ([A-Za-z0-9_-]+) from now on", msg, flags=re.IGNORECASE) + if preferred_name_match: + return set_user_preferred_name(ctx, preferred_name_match.group(1)) + + if response := _handle_memory_message(ctx, msg): + return response + + if response := _handle_reminder_message(ctx, msg): + return response + + due_response, due_context = _handle_due_reminders(ctx, msg) + if due_response is not None: + ctx.__dict__["_pending_test_context"] = due_context + return due_response + + if response := _handle_document_question(ctx, msg): + return response + + if response := _handle_custom_tool_message(ctx, msg): + return response + + if "hello" in lower_msg or "hi" in lower_msg: + return "Hello!" + + return "OK" + + +def _match_score(query: str, text: str) -> int: + stopwords = {"please", "without", "questions", "question", "there", "their", "about", "today", "going", "mention", "feeling"} + query_tokens = {token for token in _tokenize(query) if len(token) > 3 and token not in stopwords} + text_tokens = set(_tokenize(text)) + return len(query_tokens & text_tokens) + + +def _mock_relevant_context(ctx: ElroyContext, msg: str) -> list[ContextMessage]: + from elroy.repository.memories.queries import get_active_memories + + relevant_items: list[EmbeddableSqlModel] = [] + for reminder in get_active_reminders(ctx): + searchable = f"{reminder.name} {reminder.text} {reminder.reminder_context or ''}" + if _match_score(msg, searchable) > 0: + relevant_items.append(reminder) + + for memory in get_active_memories(ctx): + searchable = memory.to_fact() + if _match_score(msg, searchable) > 0: + relevant_items.append(memory) + + existing_messages = list(get_context_messages(ctx)) + unrecalled_items = [item for item in relevant_items if not any(is_in_context_message(item, msg) for msg in existing_messages)] + return to_fast_recall_tool_call(unrecalled_items) if unrecalled_items else [] + + +@tracer.chain +def process_test_message(ctx: ElroyContext, msg: str, force_tool: str | None = None) -> str: + logging.info(f"USER MESSAGE: {msg}") + if force_tool: + response = _invoke_tool_from_text(ctx, force_tool, msg) + _record_context(ctx, msg, response) + logging.info(f"ASSISTANT MESSAGE: {response}") + return response + + pending_context = list(getattr(ctx, "_pending_test_context", [])) + ctx.__dict__["_pending_test_context"] = [] + relevant_context = _mock_relevant_context(ctx, msg) + response = _answer_from_state(ctx, msg) + _record_context(ctx, msg, response, extra_messages=pending_context + relevant_context) + logging.info(f"ASSISTANT MESSAGE: {response}") + return response + + +def vector_search_by_text(ctx: ElroyContext, query: str, table: type[EmbeddableSqlModel]) -> EmbeddableSqlModel | None: + if table is Reminder: + candidates = get_active_reminders(ctx) + else: + candidates = list(ctx.db.exec(select(table).where(table.user_id == ctx.user_id)).all()) + ranked = sorted( + candidates, + key=lambda item: _match_score(query, item.to_fact()), + reverse=True, ) + return first_or_none([item for item in ranked if _match_score(query, item.to_fact()) > 0]) - bool_answer = get_boolean(full_response) - assert bool_answer == expected_answer, f"Expected {expected_answer}, got {bool_answer}. Full response: {full_response}" +def quiz_assistant_bool(expected_answer: bool, ctx: ElroyContext, question: str) -> None: + lower_question = question.lower() + reminders_summary = get_active_reminders_summary(ctx).lower() + last_response = str(getattr(ctx, "_last_test_response", "")).lower() + + if "did you just inform me about a reminder that was due" in lower_question: + bool_answer = "reminder due" in last_response or "due reminders" in last_response + elif "did the reminder i asked you to create already exist" in lower_question: + bool_answer = "already exists" in last_response or "already exist" in last_response + elif "did the reminder i asked you to delete exist" in lower_question: + bool_answer = "not found" not in last_response and "no reminder specified" not in last_response + elif "do i still have an active reminder called" in lower_question or "do i have a reminder called" in lower_question: + match = re.search(r"'([^']+)'", question) + reminder_name = match.group(1) if match else "" + bool_answer = any(reminder.name.lower() == reminder_name.lower() for reminder in get_active_reminders(ctx)) + elif "do i have any reminders about" in lower_question: + keywords = [token for token in _tokenize(lower_question) if token not in {"do", "have", "any", "reminders", "about", "or"}] + bool_answer = any(keyword in reminders_summary for keyword in keywords) + else: + raise AssertionError(f"Mock quiz handler does not understand question: {question}") + + assert bool_answer == expected_answer, f"Expected {expected_answer}, got {bool_answer}. Question: {question}" def get_active_reminders_summary(ctx: ElroyContext) -> str: