"""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