Files
cleveragents-core/features/steps/session_persistence_steps.py
T
CoreRasurae 47cad3c77b
CI / lint (pull_request) Successful in 17s
CI / typecheck (pull_request) Successful in 29s
CI / security (pull_request) Has been cancelled
CI / quality (pull_request) Has been cancelled
CI / unit_tests (pull_request) Has been cancelled
CI / integration_tests (pull_request) Has been cancelled
CI / build (pull_request) Has been cancelled
CI / docker (pull_request) Has been cancelled
CI / coverage (pull_request) Has been cancelled
feat(session): add session persistence and repositories
2026-02-18 00:10:15 +00:00

431 lines
15 KiB
Python

"""Step definitions for session persistence tests."""
from __future__ import annotations
import re
from behave import given, then, when
from behave.runner import Context
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
from cleveragents.application.services.session_service import (
PersistentSessionService,
)
from cleveragents.domain.models.core.session import (
MessageRole,
SessionImportError,
SessionNotFoundError,
)
from cleveragents.infrastructure.database.models import Base
from cleveragents.infrastructure.database.repositories import (
SessionMessageRepository,
SessionRepository,
)
# ULID pattern
_ULID_RE = re.compile(r"^[0-9A-HJKMNP-TV-Z]{26}$")
def _setup_service(context: Context) -> None:
"""Set up an in-memory database and service for testing."""
engine = create_engine("sqlite:///:memory:", echo=False)
Base.metadata.create_all(engine)
factory = sessionmaker(bind=engine, expire_on_commit=False)
# Store a single session for the whole scenario
db_session = factory()
context._db_session = db_session
context._db_engine = engine
def get_session() -> Session:
return db_session
session_repo = SessionRepository(get_session)
message_repo = SessionMessageRepository(get_session)
context.svc = PersistentSessionService(session_repo, message_repo)
context._message_repo = message_repo
# ---------------------------------------------------------------------------
# Given steps
# ---------------------------------------------------------------------------
@given("a session persistence service is initialised")
def step_session_persistence_service_initialised(context: Context) -> None:
_setup_service(context)
@given("I persist {count:d} sessions")
def step_persist_n_sessions(context: Context, count: int) -> None:
context.created_sessions = []
for _ in range(count):
s = context.svc.create()
context._db_session.commit()
context.created_sessions.append(s)
@given("I persist a new session with no actor")
def step_given_persist_session_no_actor(context: Context) -> None:
context.created_session = context.svc.create()
context._db_session.commit()
@given('I persist a new session with actor "{actor_name}"')
def step_given_persist_session_with_actor(context: Context, actor_name: str) -> None:
context.created_session = context.svc.create(actor_name=actor_name)
context._db_session.commit()
@given("I append {count:d} messages to the persisted session")
def step_given_append_n_messages(context: Context, count: int) -> None:
for i in range(count):
context.svc.append_message(
context.created_session.session_id,
MessageRole.USER,
f"Message {i}",
)
context._db_session.commit()
@given("I export the persisted session")
def step_given_export_session(context: Context) -> None:
context.export_data = context.svc.export_session(context.created_session.session_id)
# ---------------------------------------------------------------------------
# When steps
# ---------------------------------------------------------------------------
@when("I persist a new session with no actor")
def step_persist_session_no_actor(context: Context) -> None:
context.created_session = context.svc.create()
context._db_session.commit()
@when('I persist a new session with actor "{actor_name}"')
def step_persist_session_with_actor(context: Context, actor_name: str) -> None:
context.created_session = context.svc.create(actor_name=actor_name)
context._db_session.commit()
@when("I retrieve the persisted session by ID")
def step_retrieve_session_by_id(context: Context) -> None:
context.retrieved_session = context.svc.get(context.created_session.session_id)
@when("I list all persisted sessions")
def step_list_all_persisted_sessions(context: Context) -> None:
context.session_list = context.svc.list()
@when("I delete the persisted session by ID")
def step_delete_session_by_id(context: Context) -> None:
context.svc.delete(context.created_session.session_id)
context._db_session.commit()
@when("I try to delete a non-existent session")
def step_try_delete_nonexistent_session(context: Context) -> None:
try:
context.svc.delete("01NOTAVALIDULIDBUTTWENTY6")
context.delete_error = None
except SessionNotFoundError as exc:
context.delete_error = exc
@when('I append a user message "{content}" to the persisted session')
def step_append_user_message(context: Context, content: str) -> None:
context.appended_message = context.svc.append_message(
context.created_session.session_id,
MessageRole.USER,
content,
)
context._db_session.commit()
@when("I append {count:d} messages to the persisted session")
def step_append_n_messages(context: Context, count: int) -> None:
context.appended_messages = []
for i in range(count):
msg = context.svc.append_message(
context.created_session.session_id,
MessageRole.USER,
f"Message {i}",
)
context._db_session.commit()
context.appended_messages.append(msg)
@when("I try to append a message to a non-existent session")
def step_try_append_nonexistent(context: Context) -> None:
try:
context.svc.append_message(
"01NOTAVALIDULIDBUTTWENTY6",
MessageRole.USER,
"test",
)
context.append_error = None
except SessionNotFoundError as exc:
context.append_error = exc
@when("I retrieve all messages for the persisted session")
def step_retrieve_all_messages(context: Context) -> None:
context.retrieved_messages = context._message_repo.get_for_session(
context.created_session.session_id
)
@when("I retrieve messages with limit {limit:d} and offset {offset:d}")
def step_retrieve_messages_paginated(context: Context, limit: int, offset: int) -> None:
context.paginated_messages = context._message_repo.get_for_session(
context.created_session.session_id,
limit=limit,
offset=offset,
)
@when("I count messages for the persisted session")
def step_count_messages(context: Context) -> None:
context.message_count = context._message_repo.count_for_session(
context.created_session.session_id
)
@when("I export the persisted session")
def step_export_session(context: Context) -> None:
context.export_data = context.svc.export_session(context.created_session.session_id)
@when("I import the exported session data")
def step_import_exported_session(context: Context) -> None:
context.imported_session = context.svc.import_session(context.export_data)
context._db_session.commit()
@when('I try to import session data with schema version "{version}"')
def step_try_import_bad_version(context: Context, version: str) -> None:
try:
context.svc.import_session({"schema_version": version})
context.import_error = None
except SessionImportError as exc:
context.import_error = exc
@when("I tamper with the exported checksum and import")
def step_tamper_checksum_import(context: Context) -> None:
data = dict(context.export_data)
data["checksum"] = (
"0000000000000000000000000000000000000000000000000000000000000000"
)
try:
context.svc.import_session(data)
context.import_error = None
except SessionImportError as exc:
context.import_error = exc
@when("I update token usage with {inp:d} input and {out:d} output and {cost:f} cost")
def step_update_token_usage(context: Context, inp: int, out: int, cost: float) -> None:
context.svc.update_token_usage(context.created_session.session_id, inp, out, cost)
context._db_session.commit()
@when("I try to update token usage on a non-existent session")
def step_try_update_token_nonexistent(context: Context) -> None:
try:
context.svc.update_token_usage("01NOTAVALIDULIDBUTTWENTY6", 10, 5, 0.001)
context.token_error = None
except SessionNotFoundError as exc:
context.token_error = exc
@when("I persist a session with full metadata and actor")
def step_persist_session_full_metadata(context: Context) -> None:
session = context.svc.create(actor_name="local/orchestrator")
context.created_session = session
context._db_session.commit()
@when("I append a message with metadata to the persisted session")
def step_append_message_with_metadata(context: Context) -> None:
context.appended_message = context.svc.append_message(
context.created_session.session_id,
MessageRole.USER,
"metadata test",
metadata={"source": "test", "priority": 1},
)
context._db_session.commit()
# ---------------------------------------------------------------------------
# Then steps
# ---------------------------------------------------------------------------
@then("the persisted session should have a valid ULID")
def step_session_has_valid_ulid(context: Context) -> None:
assert _ULID_RE.match(context.created_session.session_id), (
f"Invalid ULID: {context.created_session.session_id}"
)
@then('the persisted session namespace should be "{ns}"')
def step_session_namespace(context: Context, ns: str) -> None:
assert context.created_session.namespace == ns
@then('the persisted session actor_name should be "{actor}"')
def step_session_actor_name(context: Context, actor: str) -> None:
assert context.created_session.actor_name == actor
@then("the retrieved session should match the created session")
def step_retrieved_matches_created(context: Context) -> None:
assert context.retrieved_session.session_id == context.created_session.session_id
assert context.retrieved_session.namespace == context.created_session.namespace
@then("the persisted session list should have {count:d} entries")
def step_session_list_count(context: Context, count: int) -> None:
assert len(context.session_list) == count, (
f"Expected {count}, got {len(context.session_list)}"
)
@then("retrieving the deleted session should raise SessionNotFoundError")
def step_retrieve_deleted_raises(context: Context) -> None:
try:
context.svc.get(context.created_session.session_id)
assert False, "Expected SessionNotFoundError" # noqa: B011
except SessionNotFoundError:
pass
@then("the session deletion should raise SessionNotFoundError")
def step_deletion_raised_error(context: Context) -> None:
assert context.delete_error is not None
assert isinstance(context.delete_error, SessionNotFoundError)
@then('the appended message role should be "{role}"')
def step_appended_role(context: Context, role: str) -> None:
assert context.appended_message.role.value == role
@then('the appended message content should be "{content}"')
def step_appended_content(context: Context, content: str) -> None:
assert context.appended_message.content == content
@then("the persisted session should have {count:d} messages with sequential ordering")
def step_session_message_count_ordered(context: Context, count: int) -> None:
messages = context._message_repo.get_for_session(context.created_session.session_id)
assert len(messages) == count
for i, msg in enumerate(messages):
assert msg.sequence == i, f"Expected seq {i}, got {msg.sequence}"
@then("the retrieved messages should be ordered by sequence")
def step_messages_ordered(context: Context) -> None:
sequences = [m.sequence for m in context.retrieved_messages]
assert sequences == sorted(sequences), f"Not ordered: {sequences}"
@then("the message append should raise SessionNotFoundError")
def step_append_raised_error(context: Context) -> None:
assert context.append_error is not None
assert isinstance(context.append_error, SessionNotFoundError)
@then("I should get exactly {count:d} messages")
def step_paginated_count(context: Context, count: int) -> None:
assert len(context.paginated_messages) == count, (
f"Expected {count}, got {len(context.paginated_messages)}"
)
@then("the first paginated message sequence should be {seq:d}")
def step_first_paginated_seq(context: Context, seq: int) -> None:
assert context.paginated_messages[0].sequence == seq
@then("the message count should be {count:d}")
def step_message_count_value(context: Context, count: int) -> None:
assert context.message_count == count
@then("the export dict should contain a schema_version key")
def step_export_has_schema_version(context: Context) -> None:
assert "schema_version" in context.export_data
@then("the export dict should contain a checksum key")
def step_export_has_checksum(context: Context) -> None:
assert "checksum" in context.export_data
@then("the export dict should contain {count:d} messages")
def step_export_message_count(context: Context, count: int) -> None:
assert len(context.export_data["messages"]) == count
@then("the imported session should have {count:d} messages")
def step_imported_message_count(context: Context, count: int) -> None:
assert len(context.imported_session.messages) == count
@then("the imported session should have a different ID from the original")
def step_imported_different_id(context: Context) -> None:
assert context.imported_session.session_id != context.created_session.session_id
@then("the import should raise SessionImportError")
def step_import_raised_error(context: Context) -> None:
assert context.import_error is not None
assert isinstance(context.import_error, SessionImportError)
@then("the persisted session token usage should show {count:d} input tokens")
def step_token_input(context: Context, count: int) -> None:
s = context.svc.get(context.created_session.session_id)
assert s.token_usage.input_tokens == count, (
f"Expected {count}, got {s.token_usage.input_tokens}"
)
@then("the persisted session token usage should show {count:d} output tokens")
def step_token_output(context: Context, count: int) -> None:
s = context.svc.get(context.created_session.session_id)
assert s.token_usage.output_tokens == count
@then("the persisted session token usage should show {cost:f} estimated cost")
def step_token_cost(context: Context, cost: float) -> None:
s = context.svc.get(context.created_session.session_id)
assert abs(s.token_usage.estimated_cost - cost) < 1e-6
@then("the token update should raise SessionNotFoundError")
def step_token_update_raised(context: Context) -> None:
assert context.token_error is not None
assert isinstance(context.token_error, SessionNotFoundError)
@then("all session fields should be preserved in the round-trip")
def step_all_fields_preserved(context: Context) -> None:
s = context.retrieved_session
assert s.session_id == context.created_session.session_id
assert s.actor_name == context.created_session.actor_name
assert s.namespace == context.created_session.namespace
@then("the retrieved message metadata should be preserved")
def step_message_metadata_preserved(context: Context) -> None:
assert len(context.retrieved_messages) >= 1
msg = context.retrieved_messages[0]
assert msg.metadata.get("source") == "test"
assert msg.metadata.get("priority") == 1