Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion elroy/io/completions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
12 changes: 6 additions & 6 deletions elroy/tools/tools_and_commands.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand Down
35 changes: 30 additions & 5 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -14,6 +14,7 @@
from elroy.db.db_manager import DbManager
from elroy.db.db_models import (
ContextMessageSet,
DocumentExcerpt,
Memory,
Message,
Reminder,
Expand All @@ -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"

Expand All @@ -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()

Expand Down Expand Up @@ -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
Expand Down
Loading