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
431 lines
15 KiB
Python
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
|