"""Step definitions for tool-add persistence (issue #621). Uses *file-based* SQLite so that opening a brand-new session factory against the same database file exercises the same code-path as the CLI (``tool add`` followed by ``tool list`` in separate process invocations). All step patterns are prefixed with ``tool-persist`` to avoid AmbiguousStep collisions with existing tool_registry_steps.py. """ from __future__ import annotations import os import tempfile from datetime import datetime from typing import Any from behave import given, then, when from behave.runner import Context from sqlalchemy import create_engine from sqlalchemy.engine import Engine from sqlalchemy.orm import Session, sessionmaker from cleveragents.infrastructure.database.models import Base from cleveragents.infrastructure.database.repositories import ( DuplicateToolError, ToolRegistryRepository, ) # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def _make_persist_tool_dict(name: str) -> dict[str, Any]: """Create a minimal valid tool dict for persistence testing.""" now_iso: str = datetime.now().isoformat() ns: str = name.split("/", 1)[0] if "/" in name else "" short: str = name.split("/", 1)[1] if "/" in name else name return { "name": name, "namespace": ns, "short_name": short, "description": f"Persist-test tool {name}", "tool_type": "tool", "source": "custom", "timeout": 300, "created_at": now_iso, "updated_at": now_iso, "resource_bindings": [], } def _file_engine(db_path: str) -> Engine: """Create a SQLAlchemy engine for a file-based SQLite DB.""" return create_engine(f"sqlite:///{db_path}", echo=False) def _file_session_factory(engine: Engine) -> sessionmaker[Session]: """Create a sessionmaker bound to *engine*.""" return sessionmaker(bind=engine, expire_on_commit=False) # --------------------------------------------------------------------------- # Given # --------------------------------------------------------------------------- @given("a tool-persist file-based SQLite database") def step_tool_persist_file_db(context: Context) -> None: """Create a temporary file-based SQLite database.""" fd, path = tempfile.mkstemp(suffix=".db", prefix="tool_persist_") os.close(fd) engine: Engine = _file_engine(path) Base.metadata.create_all(engine) engine.dispose() context.tool_persist_db_path = path @given("a tool-persist registry repository backed by the file database") def step_tool_persist_repo(context: Context) -> None: """Open a ToolRegistryRepository against the file database.""" db_path: str = context.tool_persist_db_path engine: Engine = _file_engine(db_path) factory: sessionmaker[Session] = _file_session_factory(engine) context.tool_persist_engine = engine context.tool_persist_factory = factory context.tool_persist_repo = ToolRegistryRepository(session_factory=factory) # --------------------------------------------------------------------------- # When # --------------------------------------------------------------------------- @when('a tool-persist tool named "{name}" is created') def step_tool_persist_create(context: Context, name: str) -> None: """Create a tool via the repository (should commit).""" repo: ToolRegistryRepository = context.tool_persist_repo tool_dict: dict[str, Any] = _make_persist_tool_dict(name) repo.create(tool_dict) @when('a tool-persist tool named "{name}" is created expecting a duplicate error') def step_tool_persist_create_dup(context: Context, name: str) -> None: """Attempt to create a duplicate tool and capture the error.""" repo: ToolRegistryRepository = context.tool_persist_repo tool_dict: dict[str, Any] = _make_persist_tool_dict(name) try: repo.create(tool_dict) context.tool_persist_dup_error = None except DuplicateToolError as exc: context.tool_persist_dup_error = exc @when("a fresh tool-persist registry repository is opened against the same file") def step_tool_persist_fresh_repo(context: Context) -> None: """Open a brand-new engine + session factory against the same DB file.""" # Dispose the previous engine to ensure no shared state if hasattr(context, "tool_persist_engine"): context.tool_persist_engine.dispose() db_path: str = context.tool_persist_db_path engine: Engine = _file_engine(db_path) factory: sessionmaker[Session] = _file_session_factory(engine) context.tool_persist_fresh_engine = engine context.tool_persist_fresh_factory = factory context.tool_persist_fresh_repo = ToolRegistryRepository(session_factory=factory) @when("all tools are listed via the fresh tool-persist repository") def step_tool_persist_list_fresh(context: Context) -> None: """List all tools from the fresh repository.""" repo: ToolRegistryRepository = context.tool_persist_fresh_repo context.tool_persist_listed = repo.list_all() # --------------------------------------------------------------------------- # Then # --------------------------------------------------------------------------- @then('the tool-persist list should contain "{name}"') def step_tool_persist_list_contains(context: Context, name: str) -> None: """Assert the listed tools contain a tool with the given name.""" listed: list[Any] = context.tool_persist_listed names: list[str] = [] for t in listed: if isinstance(t, dict): names.append(str(t.get("name", ""))) else: names.append(str(getattr(t, "name", ""))) assert name in names, f"Expected '{name}' in {names}" @then("the tool-persist list should have {count:d} entries") def step_tool_persist_list_count(context: Context, count: int) -> None: """Assert the number of listed tools.""" listed: list[Any] = context.tool_persist_listed assert len(listed) == count, f"Expected {count}, got {len(listed)}" @then("a tool-persist DuplicateToolError should have been raised") def step_tool_persist_dup_error(context: Context) -> None: """Assert that a DuplicateToolError was captured.""" err: DuplicateToolError | None = getattr(context, "tool_persist_dup_error", None) assert err is not None, "Expected DuplicateToolError but none was raised" assert isinstance(err, DuplicateToolError)