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
139 lines
3.9 KiB
Python
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]]()
|