feat(session): add session persistence and repositories
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
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
This commit is contained in:
@@ -0,0 +1,138 @@
|
||||
"""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]]()
|
||||
Reference in New Issue
Block a user