diff --git a/scripts/dialectic_cost_calculator.py b/scripts/dialectic_cost_calculator.py index 0c014628..af0e3275 100644 --- a/scripts/dialectic_cost_calculator.py +++ b/scripts/dialectic_cost_calculator.py @@ -170,10 +170,10 @@ def calculate_level_cost( realistic_final_answer=realistic_final, ) - model = level_config.MODEL + model = level_config.MODEL_CONFIG.model max_iterations = level_config.MAX_TOOL_ITERATIONS - thinking_budget = level_config.THINKING_BUDGET_TOKENS - provider = level_config.PROVIDER + thinking_budget = level_config.MODEL_CONFIG.thinking_budget_tokens or 0 + provider = level_config.MODEL_CONFIG.transport # Get pricing for this model pricing = MODEL_PRICING.get(model, {"input": 0, "output": 0, "cached": 0}) diff --git a/scripts/test_reasoning_levels.py b/scripts/test_reasoning_levels.py index 699b5f18..3fd38752 100755 --- a/scripts/test_reasoning_levels.py +++ b/scripts/test_reasoning_levels.py @@ -6,6 +6,7 @@ import json import os import time from datetime import datetime, timedelta, timezone +from typing import Any import httpx from dotenv import load_dotenv @@ -106,7 +107,7 @@ def load_locomo( print(f" Created session: {session_id}") # Build message batch - msg_batch = [] + msg_batch: list[dict[str, Any]] = [] for i, msg in enumerate(messages): msg_time = base_time + timedelta(seconds=i * 2) msg_batch.append( @@ -134,7 +135,7 @@ def load_locomo( def chat( client: httpx.Client, workspace_id: str, peer_id: str, query: str, level: str -) -> dict: +) -> dict[str, Any]: """Call the chat endpoint with a specific reasoning level.""" resp = client.post( f"{BASE_URL}/workspaces/{workspace_id}/peers/{peer_id}/chat", diff --git a/tests/conftest.py b/tests/conftest.py index b3697242..1ec64055 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,4 +1,8 @@ import logging +import os +import re +import time +import uuid from collections.abc import AsyncGenerator, Callable from typing import Any from unittest.mock import AsyncMock, MagicMock, patch @@ -13,7 +17,7 @@ from fastapi import Request from fastapi.responses import JSONResponse from fastapi.testclient import TestClient from nanoid import generate as generate_nanoid -from sqlalchemy import text +from sqlalchemy import create_engine, text from sqlalchemy.engine.url import URL, make_url from sqlalchemy.exc import OperationalError, ProgrammingError from sqlalchemy.ext.asyncio import ( @@ -25,7 +29,6 @@ from sqlalchemy.ext.asyncio import ( from sqlalchemy_utils import ( create_database, # pyright: ignore[reportUnknownVariableType] database_exists, # pyright: ignore[reportUnknownVariableType] - drop_database, # pyright: ignore[reportUnknownVariableType] ) from src import models @@ -128,11 +131,122 @@ def pytest_collection_modifyitems( item.add_marker(skip_live) +_RUN_ID_ENV_VAR = "HONCHO_TEST_RUN_ID" +_RUN_ID_TIME_FORMAT = "%Y%m%d%H%M%S" + +# Only a database whose name carries a run-id timestamp this old is swept. Long +# enough that no live suite is ever this stale, short enough that a leak from the +# morning is gone by the afternoon. +_STALE_DB_AGE_SECONDS = 2 * 60 * 60 + +# test_db_<14-digit timestamp>_<4 hex>[_gwN] -- only names this function minted. +# A pinned HONCHO_TEST_RUN_ID deliberately won't match, so it's never swept. +_SWEEPABLE_DB_NAME = re.compile(r"^test_db_(\d{14})_[0-9a-f]{4}(?:_gw\d+)?$") + + +def pytest_configure(config: pytest.Config) -> None: # pyright: ignore[reportUnusedParameter] + """Stamp this pytest run with an id so its databases can't collide with another run's. + + The xdist controller runs this first and its environment is inherited by the + workers it spawns, so `setdefault` gives every worker in a run the same id + while separate runs (concurrent worktrees, two agents, a local run alongside + CI) each get their own. Set the env var yourself to pin a stable name. + + The id leads with a sortable local-time timestamp so leaked databases can be + aged out (see `_sweep_stale_test_databases`); the random tail keeps two runs + starting in the same second apart. + """ + + os.environ.setdefault( + _RUN_ID_ENV_VAR, + f"{time.strftime(_RUN_ID_TIME_FORMAT)}_{uuid.uuid4().hex[:4]}", + ) + + # Workers inherit the controller's env and would each redo this. + if os.environ.get("PYTEST_XDIST_WORKER") is None: + _sweep_stale_test_databases() + + def _get_test_db_url(worker_id: str) -> URL: """Get a worker-specific test database URL for pytest-xdist parallelism.""" - db_name = "test_db" if worker_id == "master" else f"test_db_{worker_id}" - return CONNECTION_URI.set(database=db_name) + run_id = os.environ.get(_RUN_ID_ENV_VAR, "local") + suffix = "" if worker_id == "master" else f"_{worker_id}" + return CONNECTION_URI.set(database=f"test_db_{run_id}{suffix}") + + +def _drop_database(db_url: URL) -> None: + """Drop a test database, evicting any connections still holding it open. + + WITH (FORCE) (pg13+) is what makes this reliable: a pooled connection that + outlives engine disposal, or an xdist worker killed mid-query, otherwise + leaves the drop failing with "database is being accessed by other users". + """ + + name = db_url.database + if not name: + return + + # Maintenance connection: you cannot drop the database you're connected to. + engine = create_engine( + db_url.set(database="postgres"), isolation_level="AUTOCOMMIT" + ) + try: + with engine.connect() as conn: + conn.exec_driver_sql(f'DROP DATABASE IF EXISTS "{name}" WITH (FORCE)') + finally: + engine.dispose() + + +def _sweep_stale_test_databases() -> None: + """Reclaim test databases left behind by runs that died before teardown. + + A run killed by SIGKILL, an IDE stop button, an OOM'd worker or `-x` on a hang + never reaches the `db_engine` teardown, and since every run mints its own + database name nothing later reuses (and thus cleans) it. + + Two guards keep this from touching a suite that is currently running, which is + the whole point of per-run names: + + - the run-id timestamp in the name must be older than `_STALE_DB_AGE_SECONDS` + - the database must have no backends connected to it right now + + Each covers the other's blind spot: the age check is immune to the race where + a database has been created but its first worker hasn't connected yet, and the + connection check catches a genuinely long-running suite. Failure to sweep is + logged and ignored -- it must never fail a test session. + """ + + cutoff = time.strftime( + _RUN_ID_TIME_FORMAT, time.localtime(time.time() - _STALE_DB_AGE_SECONDS) + ) + + try: + engine = create_engine( + CONNECTION_URI.set(database="postgres"), isolation_level="AUTOCOMMIT" + ) + try: + with engine.connect() as conn: + names = [ + row[0] + for row in conn.exec_driver_sql( + "SELECT datname FROM pg_database d " + + "WHERE NOT EXISTS (" + + " SELECT 1 FROM pg_stat_activity WHERE datname = d.datname" + + ")" + ) + ] + finally: + engine.dispose() + + for name in names: + match = _SWEEPABLE_DB_NAME.match(name) + if match is None or match.group(1) >= cutoff: + continue + logger.info(f"Dropping stale test database: {name}") + _drop_database(CONNECTION_URI.set(database=name)) + except Exception as e: + logger.warning(f"Could not sweep stale test databases: {e}") # Test API authorization - no longer needed as module-level constants @@ -193,11 +307,21 @@ async def setup_test_database(db_url: URL): return engine -async def _truncate_all_tables(engine: AsyncEngine) -> None: - """Remove all data from every mapped table while resetting identities.""" +async def _clear_all_tables(engine: AsyncEngine) -> None: + """Remove all data from every mapped table between tests. + + Uses DELETE rather than TRUNCATE: TRUNCATE rewrites the relfilenode of every + table and index it touches, so it costs a flat ~33ms for this schema's 11 + tables / 41 indexes no matter how few rows a test actually wrote. DELETE of + the same (near-empty) tables, batched into one round trip, is ~3ms. Tables go + in reverse dependency order so foreign keys are satisfied without CASCADE. + + This does not reset identity sequences, so tests must not assert on absolute + generated id values -- compare against the ids the test itself created. + """ table_names: list[str] = [] - for table in Base.metadata.sorted_tables: + for table in reversed(Base.metadata.sorted_tables): if table.schema: table_names.append(f'"{table.schema}"."{table.name}"') else: @@ -206,9 +330,9 @@ async def _truncate_all_tables(engine: AsyncEngine) -> None: if not table_names: return - joined_names = ", ".join(table_names) + statement = "; ".join(f"DELETE FROM {name}" for name in table_names) async with engine.begin() as conn: - await conn.execute(text(f"TRUNCATE {joined_names} RESTART IDENTITY CASCADE")) + await conn.exec_driver_sql(statement) @pytest_asyncio.fixture(scope="session") @@ -242,7 +366,7 @@ async def db_engine(worker_id: str): for table in Base.metadata.tables.values(): table.schema = original_schema - drop_database(test_db_url) + _drop_database(test_db_url) @pytest_asyncio.fixture(scope="function") @@ -256,7 +380,7 @@ async def db_session(db_engine: AsyncEngine): finally: await session.rollback() finally: - await _truncate_all_tables(db_engine) + await _clear_all_tables(db_engine) @pytest_asyncio.fixture(scope="session") @@ -430,7 +554,7 @@ async def sample_data( db_session.add(test_peer) # Commit so data is visible to independent tracked_db sessions. - # _truncate_all_tables handles cleanup between tests. + # _clear_all_tables handles cleanup between tests. await db_session.commit() yield test_workspace, test_peer diff --git a/tests/utils/test_agent_tools.py b/tests/utils/test_agent_tools.py index 71b93eb0..c13fb1e7 100644 --- a/tests/utils/test_agent_tools.py +++ b/tests/utils/test_agent_tools.py @@ -129,7 +129,7 @@ async def tool_test_data( # Commit so data is visible to independent tracked_db sessions. # Tool handlers no longer share the test's db_session — they open # their own short-lived sessions via tracked_db. - # _truncate_all_tables handles cleanup between tests. + # _clear_all_tables handles cleanup between tests. await db_session.commit() yield workspace, peer1, peer2, session, messages, documents