Files
cleveragents-core/robot/helper_session_persistence.py
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

139 lines
3.9 KiB
Python

"""Helper script for Robot Framework session persistence smoke tests.
Usage:
python robot/helper_session_persistence.py <test-name>
Supported test names:
session-round-trip Create/read/delete session round-trip
message-append Append messages and verify ordering
export-import Export and import a session
token-usage Update and verify token usage
"""
from __future__ import annotations
import sys
from pathlib import Path
# Ensure src is on the path
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from cleveragents.application.services.session_service import (
PersistentSessionService,
)
from cleveragents.domain.models.core.session import MessageRole
from cleveragents.infrastructure.database.models import Base
from cleveragents.infrastructure.database.repositories import (
SessionMessageRepository,
SessionRepository,
)
def _make_service() -> tuple:
"""Create in-memory DB and service."""
engine = create_engine("sqlite:///:memory:", echo=False)
Base.metadata.create_all(engine)
factory = sessionmaker(bind=engine, expire_on_commit=False)
db_session = factory()
def get_session():
return db_session
session_repo = SessionRepository(get_session)
message_repo = SessionMessageRepository(get_session)
svc = PersistentSessionService(session_repo, message_repo)
return svc, db_session
def test_session_round_trip() -> None:
svc, db_session = _make_service()
# Create
session = svc.create(actor_name="local/orchestrator")
db_session.commit()
# Read
fetched = svc.get(session.session_id)
assert fetched.session_id == session.session_id
assert fetched.actor_name == "local/orchestrator"
# Delete
svc.delete(session.session_id)
db_session.commit()
try:
svc.get(session.session_id)
raise AssertionError("Expected SessionNotFoundError")
except Exception:
pass
print("session-round-trip-ok")
def test_message_append() -> None:
svc, db_session = _make_service()
session = svc.create()
db_session.commit()
for i in range(5):
svc.append_message(session.session_id, MessageRole.USER, f"msg-{i}")
db_session.commit()
# Verify message count via export (messages are loaded)
export = svc.export_session(session.session_id)
assert len(export["messages"]) == 5
print("message-append-ok")
def test_export_import() -> None:
svc, db_session = _make_service()
session = svc.create(actor_name="local/test-actor")
db_session.commit()
svc.append_message(session.session_id, MessageRole.USER, "hello")
svc.append_message(session.session_id, MessageRole.ASSISTANT, "world")
db_session.commit()
export = svc.export_session(session.session_id)
assert export["schema_version"] == "1.0"
assert "checksum" in export
imported = svc.import_session(export)
db_session.commit()
assert imported.session_id != session.session_id
assert len(imported.messages) == 2
print("export-import-ok")
def test_token_usage() -> None:
svc, db_session = _make_service()
session = svc.create()
db_session.commit()
svc.update_token_usage(session.session_id, 100, 50, 0.005)
db_session.commit()
fetched = svc.get(session.session_id)
assert fetched.token_usage.input_tokens == 100
assert fetched.token_usage.output_tokens == 50
assert abs(fetched.token_usage.estimated_cost - 0.005) < 1e-6
print("token-usage-ok")
TESTS = {
"session-round-trip": test_session_round_trip,
"message-append": test_message_append,
"export-import": test_export_import,
"token-usage": test_token_usage,
}
if __name__ == "__main__":
if len(sys.argv) < 2 or sys.argv[1] not in TESTS:
print(f"Usage: {sys.argv[0]} <{'|'.join(TESTS)}>", file=sys.stderr)
sys.exit(2)
TESTS[sys.argv[1]]()