612 lines
23 KiB
Python
612 lines
23 KiB
Python
"""Step definitions for Session domain model tests."""
|
|
|
|
from typing import Any
|
|
|
|
from behave import given, then, when
|
|
from behave.runner import Context
|
|
from pydantic import ValidationError
|
|
from ulid import ULID
|
|
|
|
from cleveragents.domain.models.core.session import (
|
|
MessageRole,
|
|
Session,
|
|
SessionExportError,
|
|
SessionImportError,
|
|
SessionMessage,
|
|
SessionNotFoundError,
|
|
SessionServiceError,
|
|
SessionTokenUsage,
|
|
)
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
_VALID_ULID = str(ULID())
|
|
|
|
|
|
def _make_session(**overrides: Any) -> Session:
|
|
"""Create a Session with sensible defaults, allowing overrides."""
|
|
defaults: dict[str, Any] = {
|
|
"session_id": str(ULID()),
|
|
}
|
|
defaults.update(overrides)
|
|
return Session(**defaults)
|
|
|
|
|
|
def _make_message(**overrides: Any) -> SessionMessage:
|
|
"""Create a SessionMessage with sensible defaults."""
|
|
defaults: dict[str, Any] = {
|
|
"message_id": str(ULID()),
|
|
"role": MessageRole.USER,
|
|
"content": "Test message",
|
|
"sequence": 0,
|
|
}
|
|
defaults.update(overrides)
|
|
return SessionMessage(**defaults)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Session Creation Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I create a session with a valid ULID")
|
|
def session_model_create_with_valid_ulid(context: Context) -> None:
|
|
"""Create a session with a valid ULID."""
|
|
context.session_model = _make_session()
|
|
context.session_model_error = None
|
|
|
|
|
|
@when('I create a session with actor name "{actor_name}"')
|
|
def session_model_create_with_actor(context: Context, actor_name: str) -> None:
|
|
"""Create a session with a specific actor name."""
|
|
context.session_model = _make_session(actor_name=actor_name)
|
|
context.session_model_error = None
|
|
|
|
|
|
@when("I create a session without actor name")
|
|
def session_model_create_without_actor(context: Context) -> None:
|
|
"""Create a session without an actor name."""
|
|
context.session_model = _make_session(actor_name=None)
|
|
context.session_model_error = None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Session Validation Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when('I try to create a session with invalid ID "{session_id}"')
|
|
def session_model_try_invalid_id(context: Context, session_id: str) -> None:
|
|
"""Attempt to create a session with an invalid ID."""
|
|
context.session_model_error = None
|
|
context.session_model = None
|
|
try:
|
|
context.session_model = Session(session_id=session_id)
|
|
except ValidationError as e:
|
|
context.session_model_error = e
|
|
|
|
|
|
@when('I try to create a session with invalid actor name "{actor_name}"')
|
|
def session_model_try_invalid_actor(context: Context, actor_name: str) -> None:
|
|
"""Attempt to create a session with an invalid actor name."""
|
|
context.session_model_error = None
|
|
context.session_model = None
|
|
try:
|
|
context.session_model = _make_session(actor_name=actor_name)
|
|
except ValidationError as e:
|
|
context.session_model_error = e
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Session Assertions
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@then("the session model should be created")
|
|
def session_model_check_created(context: Context) -> None:
|
|
"""Verify the session was created."""
|
|
assert context.session_model is not None, "Session should be created"
|
|
|
|
|
|
@then("the session model session_id should be a valid ULID")
|
|
def session_model_check_ulid(context: Context) -> None:
|
|
"""Verify the session_id is a valid ULID."""
|
|
import re
|
|
|
|
pattern = r"^[0-9A-HJKMNP-TV-Z]{26}$"
|
|
assert re.match(pattern, context.session_model.session_id), (
|
|
f"session_id '{context.session_model.session_id}' is not a valid ULID"
|
|
)
|
|
|
|
|
|
@then('the session model actor_name should be "{expected}"')
|
|
def session_model_check_actor_name(context: Context, expected: str) -> None:
|
|
"""Check the session actor_name."""
|
|
assert context.session_model.actor_name == expected, (
|
|
f"Expected actor_name '{expected}', got '{context.session_model.actor_name}'"
|
|
)
|
|
|
|
|
|
@then("the session model actor_name should be empty")
|
|
def session_model_check_actor_name_empty(context: Context) -> None:
|
|
"""Check the session actor_name is None."""
|
|
assert context.session_model.actor_name is None, (
|
|
f"Expected actor_name to be None, got '{context.session_model.actor_name}'"
|
|
)
|
|
|
|
|
|
@then('the session model namespace should be "{expected}"')
|
|
def session_model_check_namespace(context: Context, expected: str) -> None:
|
|
"""Check the session namespace."""
|
|
assert context.session_model.namespace == expected, (
|
|
f"Expected namespace '{expected}', got '{context.session_model.namespace}'"
|
|
)
|
|
|
|
|
|
@then("a session model validation error should be raised")
|
|
def session_model_check_validation_error(context: Context) -> None:
|
|
"""Verify that a validation error was raised."""
|
|
assert context.session_model_error is not None, (
|
|
"Expected a validation error to be raised"
|
|
)
|
|
|
|
|
|
@then('the session model error should mention "{text}"')
|
|
def session_model_check_error_message(context: Context, text: str) -> None:
|
|
"""Check the error message contains expected text."""
|
|
error_str = str(context.session_model_error)
|
|
assert text.lower() in error_str.lower(), (
|
|
f"Expected error to mention '{text}', got: {error_str}"
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Message Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given("a session with no messages")
|
|
def session_model_given_empty_session(context: Context) -> None:
|
|
"""Create a session with no messages."""
|
|
context.session_model = _make_session()
|
|
context.session_model_error = None
|
|
|
|
|
|
@when('I append a user message "{content}"')
|
|
def session_model_append_user_message(context: Context, content: str) -> None:
|
|
"""Append a user message to the session."""
|
|
context.session_model.append_message(MessageRole.USER, content)
|
|
|
|
|
|
@when('I append an assistant message "{content}"')
|
|
def session_model_append_assistant_message(context: Context, content: str) -> None:
|
|
"""Append an assistant message to the session."""
|
|
context.session_model.append_message(MessageRole.ASSISTANT, content)
|
|
|
|
|
|
@then("the session model should have {count:d} message")
|
|
def session_model_check_message_count_singular(context: Context, count: int) -> None:
|
|
"""Check message count (singular)."""
|
|
actual = context.session_model.message_count
|
|
assert actual == count, f"Expected {count} message(s), got {actual}"
|
|
|
|
|
|
@then("the session model should have {count:d} messages")
|
|
def session_model_check_message_count_plural(context: Context, count: int) -> None:
|
|
"""Check message count (plural)."""
|
|
actual = context.session_model.message_count
|
|
assert actual == count, f"Expected {count} message(s), got {actual}"
|
|
|
|
|
|
@then('the session last message content should be "{expected}"')
|
|
def session_model_check_last_message_content(context: Context, expected: str) -> None:
|
|
"""Check the last message content."""
|
|
last = context.session_model.last_message
|
|
assert last is not None, "Expected a last message"
|
|
assert last.content == expected, (
|
|
f"Expected content '{expected}', got '{last.content}'"
|
|
)
|
|
|
|
|
|
@then('the session last message role should be "{expected}"')
|
|
def session_model_check_last_message_role(context: Context, expected: str) -> None:
|
|
"""Check the last message role."""
|
|
last = context.session_model.last_message
|
|
assert last is not None, "Expected a last message"
|
|
assert last.role.value == expected, (
|
|
f"Expected role '{expected}', got '{last.role.value}'"
|
|
)
|
|
|
|
|
|
@then("the session last message sequence should be {expected:d}")
|
|
def session_model_check_last_message_sequence(context: Context, expected: int) -> None:
|
|
"""Check the last message sequence."""
|
|
last = context.session_model.last_message
|
|
assert last is not None, "Expected a last message"
|
|
assert last.sequence == expected, (
|
|
f"Expected sequence {expected}, got {last.sequence}"
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Tool Message Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I try to create a tool message without tool_call_id")
|
|
def session_model_try_tool_message_no_call_id(context: Context) -> None:
|
|
"""Attempt to create a tool message without tool_call_id."""
|
|
context.session_model_error = None
|
|
try:
|
|
_make_message(role=MessageRole.TOOL, content="Result", tool_call_id=None)
|
|
except ValidationError as e:
|
|
context.session_model_error = e
|
|
|
|
|
|
@when('I create a tool message with tool_call_id "{call_id}"')
|
|
def session_model_create_tool_message(context: Context, call_id: str) -> None:
|
|
"""Create a tool message with tool_call_id."""
|
|
context.session_message = _make_message(
|
|
role=MessageRole.TOOL, content="Result", tool_call_id=call_id
|
|
)
|
|
context.session_model_error = None
|
|
|
|
|
|
@when('I create a user message "{content}"')
|
|
def session_model_create_user_message(context: Context, content: str) -> None:
|
|
"""Create a user message."""
|
|
context.session_message = _make_message(role=MessageRole.USER, content=content)
|
|
context.session_model_error = None
|
|
|
|
|
|
@then("the session message should be created")
|
|
def session_model_check_message_created(context: Context) -> None:
|
|
"""Verify the message was created."""
|
|
assert context.session_message is not None, "Message should be created"
|
|
|
|
|
|
@then('the session message tool_call_id should be "{expected}"')
|
|
def session_model_check_message_tool_call_id(context: Context, expected: str) -> None:
|
|
"""Check the message tool_call_id."""
|
|
assert context.session_message.tool_call_id == expected, (
|
|
f"Expected tool_call_id '{expected}', "
|
|
f"got '{context.session_message.tool_call_id}'"
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Content Validation Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I try to create a message with whitespace-only content")
|
|
def session_model_try_whitespace_content(context: Context) -> None:
|
|
"""Attempt to create a message with whitespace-only content."""
|
|
context.session_model_error = None
|
|
try:
|
|
_make_message(content=" \t\n ")
|
|
except ValidationError as e:
|
|
context.session_model_error = e
|
|
|
|
|
|
@when("I try to create a message with empty content")
|
|
def session_model_try_empty_content(context: Context) -> None:
|
|
"""Attempt to create a message with empty content."""
|
|
context.session_model_error = None
|
|
try:
|
|
_make_message(content="")
|
|
except ValidationError as e:
|
|
context.session_model_error = e
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Message Ordering Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I try to create a session with out-of-order messages")
|
|
def session_model_try_out_of_order_messages(context: Context) -> None:
|
|
"""Attempt to create a session with out-of-order messages."""
|
|
context.session_model_error = None
|
|
context.session_model = None
|
|
msg1 = _make_message(sequence=1, content="Second")
|
|
msg2 = _make_message(sequence=0, content="First")
|
|
try:
|
|
context.session_model = _make_session(messages=[msg1, msg2])
|
|
except ValidationError as e:
|
|
context.session_model_error = e
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Plan Linking Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when('I link plan "{plan_id}" to the session')
|
|
def session_model_link_plan(context: Context, plan_id: str) -> None:
|
|
"""Link a plan to the session."""
|
|
context.session_model.link_plan(plan_id)
|
|
|
|
|
|
@then("the session should have {count:d} linked plan")
|
|
def session_model_check_linked_plan_count_singular(
|
|
context: Context, count: int
|
|
) -> None:
|
|
"""Check linked plan count (singular)."""
|
|
actual = len(context.session_model.linked_plan_ids)
|
|
assert actual == count, f"Expected {count} linked plan(s), got {actual}"
|
|
|
|
|
|
@then("the session should have {count:d} linked plans")
|
|
def session_model_check_linked_plan_count_plural(context: Context, count: int) -> None:
|
|
"""Check linked plan count (plural)."""
|
|
actual = len(context.session_model.linked_plan_ids)
|
|
assert actual == count, f"Expected {count} linked plan(s), got {actual}"
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Token Usage Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@then("the session token usage input_tokens should be {expected:d}")
|
|
def session_model_check_token_input(context: Context, expected: int) -> None:
|
|
"""Check token usage input_tokens."""
|
|
actual = context.session_model.token_usage.input_tokens
|
|
assert actual == expected, f"Expected input_tokens {expected}, got {actual}"
|
|
|
|
|
|
@then("the session token usage output_tokens should be {expected:d}")
|
|
def session_model_check_token_output(context: Context, expected: int) -> None:
|
|
"""Check token usage output_tokens."""
|
|
actual = context.session_model.token_usage.output_tokens
|
|
assert actual == expected, f"Expected output_tokens {expected}, got {actual}"
|
|
|
|
|
|
@then("the session token usage estimated_cost should be {expected:g}")
|
|
def session_model_check_token_cost(context: Context, expected: float) -> None:
|
|
"""Check token usage estimated_cost."""
|
|
actual = context.session_model.token_usage.estimated_cost
|
|
assert abs(actual - expected) < 1e-9, (
|
|
f"Expected estimated_cost {expected}, got {actual}"
|
|
)
|
|
|
|
|
|
@when(
|
|
"I create a session with token usage {input_t:d} input"
|
|
" {output_t:d} output {cost:g} cost"
|
|
)
|
|
def session_model_create_with_token_usage(
|
|
context: Context, input_t: int, output_t: int, cost: float
|
|
) -> None:
|
|
"""Create a session with specific token usage."""
|
|
context.session_model = _make_session(
|
|
token_usage=SessionTokenUsage(
|
|
input_tokens=input_t,
|
|
output_tokens=output_t,
|
|
estimated_cost=cost,
|
|
)
|
|
)
|
|
context.session_model_error = None
|
|
|
|
|
|
@when("I try to create token usage with negative input tokens")
|
|
def session_model_try_negative_tokens(context: Context) -> None:
|
|
"""Attempt to create token usage with negative input tokens."""
|
|
context.session_model_error = None
|
|
try:
|
|
SessionTokenUsage(input_tokens=-1)
|
|
except ValidationError as e:
|
|
context.session_model_error = e
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Export Dict Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given("a session with some messages")
|
|
def session_model_given_session_with_messages(context: Context) -> None:
|
|
"""Create a session with some messages."""
|
|
context.session_model = _make_session()
|
|
context.session_model.append_message(MessageRole.USER, "Hello")
|
|
context.session_model.append_message(MessageRole.ASSISTANT, "Hi there")
|
|
context.session_model_error = None
|
|
|
|
|
|
@when("I export the session")
|
|
def session_model_export(context: Context) -> None:
|
|
"""Export the session to a dict."""
|
|
context.session_export_dict = context.session_model.as_export_dict()
|
|
|
|
|
|
@then('the export dict should have key "{key}"')
|
|
def session_model_check_export_key(context: Context, key: str) -> None:
|
|
"""Check the export dict has a specific key."""
|
|
assert key in context.session_export_dict, (
|
|
f"Expected key '{key}' in export dict, "
|
|
f"keys: {list(context.session_export_dict)}"
|
|
)
|
|
|
|
|
|
@then("the export dict messages count should match session message count")
|
|
def session_model_check_export_messages_count(context: Context) -> None:
|
|
"""Check export dict messages count matches session."""
|
|
export_count = len(context.session_export_dict["messages"])
|
|
session_count = context.session_model.message_count
|
|
assert export_count == session_count, (
|
|
f"Export has {export_count} messages, session has {session_count}"
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# CLI Dict Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I get the session CLI dict")
|
|
def session_model_get_cli_dict(context: Context) -> None:
|
|
"""Get the CLI dict for the session."""
|
|
context.session_cli_dict = context.session_model.as_cli_dict()
|
|
|
|
|
|
@then('the session cli dict should have key "{key}"')
|
|
def session_model_check_cli_dict_key(context: Context, key: str) -> None:
|
|
"""Check the CLI dict has a specific key."""
|
|
assert key in context.session_cli_dict, (
|
|
f"Expected key '{key}' in cli dict, keys: {list(context.session_cli_dict)}"
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Automation Level Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given('a session with automation level "{level}"')
|
|
def session_model_given_with_automation(context: Context, level: str) -> None:
|
|
"""Create a session (automation_level field was removed)."""
|
|
context.session_model = _make_session()
|
|
context.session_model.append_message(MessageRole.USER, "Test")
|
|
context.session_model_error = None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Empty Session Property Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@then("the session message count should be {expected:d}")
|
|
def session_model_check_msg_count_property(context: Context, expected: int) -> None:
|
|
"""Check message_count property."""
|
|
actual = context.session_model.message_count
|
|
assert actual == expected, f"Expected {expected}, got {actual}"
|
|
|
|
|
|
@then("the session is_empty should be true")
|
|
def session_model_check_is_empty_true(context: Context) -> None:
|
|
"""Check is_empty is True."""
|
|
assert context.session_model.is_empty is True, "Expected session to be empty"
|
|
|
|
|
|
@then("the session last_message should be none")
|
|
def session_model_check_last_message_none(context: Context) -> None:
|
|
"""Check last_message is None."""
|
|
assert context.session_model.last_message is None, (
|
|
"Expected last_message to be None"
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# get_messages Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given("a session with {count:d} messages")
|
|
def session_model_given_n_messages(context: Context, count: int) -> None:
|
|
"""Create a session with N messages."""
|
|
context.session_model = _make_session()
|
|
for i in range(count):
|
|
context.session_model.append_message(MessageRole.USER, f"Message {i}")
|
|
context.session_model_error = None
|
|
|
|
|
|
@when("I get messages with no limit")
|
|
def session_model_get_messages_no_limit(context: Context) -> None:
|
|
"""Get all messages."""
|
|
context.session_messages_result = context.session_model.get_messages()
|
|
|
|
|
|
@when("I get messages with limit {limit:d}")
|
|
def session_model_get_messages_with_limit(context: Context, limit: int) -> None:
|
|
"""Get messages with a limit."""
|
|
context.session_messages_result = context.session_model.get_messages(limit=limit)
|
|
|
|
|
|
@when("I get messages with offset {offset:d}")
|
|
def session_model_get_messages_with_offset(context: Context, offset: int) -> None:
|
|
"""Get messages with an offset."""
|
|
context.session_messages_result = context.session_model.get_messages(offset=offset)
|
|
|
|
|
|
@when("I get messages with limit {limit:d} and offset {offset:d}")
|
|
def session_model_get_messages_with_limit_offset(
|
|
context: Context, limit: int, offset: int
|
|
) -> None:
|
|
"""Get messages with limit and offset."""
|
|
context.session_messages_result = context.session_model.get_messages(
|
|
limit=limit, offset=offset
|
|
)
|
|
|
|
|
|
@then("I should get {count:d} messages")
|
|
def session_model_check_messages_result_count(context: Context, count: int) -> None:
|
|
"""Check the number of messages returned."""
|
|
actual = len(context.session_messages_result)
|
|
assert actual == count, f"Expected {count} messages, got {actual}"
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Error Type Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@then("SessionServiceError should be an Exception subclass")
|
|
def session_model_check_service_error_base(context: Context) -> None:
|
|
"""Verify SessionServiceError inherits from Exception."""
|
|
assert issubclass(SessionServiceError, Exception)
|
|
|
|
|
|
@then("SessionNotFoundError should be a SessionServiceError subclass")
|
|
def session_model_check_not_found_error(context: Context) -> None:
|
|
"""Verify SessionNotFoundError inherits from SessionServiceError."""
|
|
assert issubclass(SessionNotFoundError, SessionServiceError)
|
|
|
|
|
|
@then("SessionExportError should be a SessionServiceError subclass")
|
|
def session_model_check_export_error(context: Context) -> None:
|
|
"""Verify SessionExportError inherits from SessionServiceError."""
|
|
assert issubclass(SessionExportError, SessionServiceError)
|
|
|
|
|
|
@then("SessionImportError should be a SessionServiceError subclass")
|
|
def session_model_check_import_error(context: Context) -> None:
|
|
"""Verify SessionImportError inherits from SessionServiceError."""
|
|
assert issubclass(SessionImportError, SessionServiceError)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# MessageRole Enum Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@then('the MessageRole enum should have values "{values}"')
|
|
def session_model_check_message_role_enum(context: Context, values: str) -> None:
|
|
"""Verify MessageRole enum values."""
|
|
expected = set(values.split(","))
|
|
actual = {e.value for e in MessageRole}
|
|
assert actual == expected, f"Expected {expected}, got {actual}"
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Metadata Steps
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I create a session with metadata")
|
|
def session_model_create_with_metadata(context: Context) -> None:
|
|
"""Create a session with metadata."""
|
|
context.session_model = _make_session(metadata={"env": "test", "version": "1.0"})
|
|
context.session_model_error = None
|
|
|
|
|
|
@then('the session metadata should contain key "{key}"')
|
|
def session_model_check_metadata_key(context: Context, key: str) -> None:
|
|
"""Check metadata contains a key."""
|
|
assert key in context.session_model.metadata, (
|
|
f"Expected key '{key}' in metadata, "
|
|
f"keys: {list(context.session_model.metadata)}"
|
|
)
|