diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py new file mode 100644 index 0000000..b63f213 --- /dev/null +++ b/backend/tests/conftest.py @@ -0,0 +1,57 @@ +"""Shared setup for the test suite. + +pytest imports this file before it imports any test module, which is the only +reason the database redirection below works. `app.database` reads +`AIDND_DB_PATH` at import and builds `engine` from it once, so the variable has +to be set before the first `from app...` line anywhere in the suite. + +Every test module used to carry its own copy of that redirection. Only the first +one to be imported ever took effect, because the engine already existed by the +time the second one ran. The other copies created a temp file that nothing +opened and nothing deleted. One copy here does the job, and it cleans up after +itself. + +The tests share one database. That is not new: they already did. Each `client` +fixture calls `Base.metadata.create_all` on setup and `drop_all` on teardown, so +no test sees another test's rows. +""" +import os +import tempfile + +_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) +_tmp.close() +os.environ["AIDND_DB_PATH"] = _tmp.name +# A real `AIDND_DATABASE_URL` or `DATABASE_URL` in the developer's shell points +# at Postgres, and `app.database` prefers either over the SQLite path above. +# Clear both, so running the suite never touches a server database. +os.environ.pop("AIDND_DATABASE_URL", None) +os.environ.pop("DATABASE_URL", None) + +import pytest # noqa: E402 Import order is load-bearing; see above. + +from fakes import ScriptedProvider # noqa: E402 + + +@pytest.fixture(autouse=True) +def reset_scripted_provider(): + """Clears the fake provider's state between tests. + + `ScriptedProvider` keeps its replies and its call count on the class, because + the code under test constructs the provider itself and a test cannot reach + the instance. Class state outlives a test, so reset it here rather than + trusting every fixture to remember. + """ + ScriptedProvider.replies = [] + ScriptedProvider.calls = 0 + ScriptedProvider.prompts = [] + yield + + +def pytest_sessionfinish(session, exitstatus): + """Deletes the temporary database once the run ends.""" + try: + os.unlink(_tmp.name) + except OSError: + # The file is already gone, or Windows still holds a handle on it. It is + # in the temp directory either way, so leaving it costs nothing. + pass diff --git a/backend/tests/fakes.py b/backend/tests/fakes.py new file mode 100644 index 0000000..8b035c9 --- /dev/null +++ b/backend/tests/fakes.py @@ -0,0 +1,41 @@ +"""Stand-ins for the parts of the app a test must not really call. + +Import these rather than writing another copy. Nine test modules each carried +their own `ScriptedProvider`, and the copies had drifted into four different +feature sets, so a test that needed to raise a provider error had to be written +in one of the files whose copy supported it. +""" + + +class ScriptedProvider: + """Streams canned replies in place of `OpenAICompatibleProvider`. + + Set `replies` to the texts the model returns, one per call. The last entry + repeats once the list runs out, so a test that plays more turns than it + scripted still gets text. To drive the provider-error path, put an + `Exception` in the list. It is raised rather than streamed. + + State lives on the class, not on the instance, because the turn engine + constructs the provider itself and a test never sees the object. The autouse + `reset_scripted_provider` fixture in `conftest.py` clears it between tests. + + `prompts` records every assembled `(system, story)` pair, which is what a + test asserts on to check what the model was shown. + """ + + last_usage = None + replies: list = [] + calls = 0 + prompts: list = [] + + def __init__(self, *a, **k): + pass + + async def generate(self, parts, *, temperature, max_tokens): + index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1) + ScriptedProvider.calls += 1 + ScriptedProvider.prompts.append((parts.system, parts.story)) + reply = ScriptedProvider.replies[index] + if isinstance(reply, Exception): + raise reply + yield ("text", reply) diff --git a/backend/tests/test_accesslog.py b/backend/tests/test_accesslog.py index d014974..5a1da0c 100644 --- a/backend/tests/test_accesslog.py +++ b/backend/tests/test_accesslog.py @@ -9,15 +9,6 @@ accounts on a schedule, and a log that deletes itself is not a log. python -m pytest tests/test_accesslog.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient diff --git a/backend/tests/test_action_paging.py b/backend/tests/test_action_paging.py index 6f4280a..4c09d7c 100644 --- a/backend/tests/test_action_paging.py +++ b/backend/tests/test_action_paging.py @@ -12,15 +12,6 @@ to be scrolling. An anchor means the same thing before and after. python -m pytest tests/test_action_paging.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient diff --git a/backend/tests/test_analytics.py b/backend/tests/test_analytics.py index e8765c4..f0e3eb4 100644 --- a/backend/tests/test_analytics.py +++ b/backend/tests/test_analytics.py @@ -9,15 +9,6 @@ the dashboard, and cannot inflate what it reports beyond hitting the page. python -m pytest tests/test_analytics.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - from datetime import timedelta import pytest diff --git a/backend/tests/test_attempt_siblings.py b/backend/tests/test_attempt_siblings.py index dbe0252..2afbbb8 100644 --- a/backend/tests/test_attempt_siblings.py +++ b/backend/tests/test_attempt_siblings.py @@ -8,15 +8,6 @@ story, and the arrangement costs neither an extra prompt nor an extra turn. python -m pytest tests/test_attempt_siblings.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient @@ -26,9 +17,10 @@ from app import attempts, auth, limits, models, tree from app.context import cursors, history from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts from app.routers import adventures +from fakes import ScriptedProvider + SCHEMA = {"player": {"hp": {"min": 0, "max": 100, "initial": 100}}} GOLD_SCRIPT = """ @@ -40,22 +32,6 @@ modifier(text); """ -class ScriptedProvider: - last_usage = None - replies: list = [] - calls = 0 - prompts: list = [] - - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1) - ScriptedProvider.calls += 1 - ScriptedProvider.prompts.append((parts.system, parts.story)) - yield ("text", ScriptedProvider.replies[index]) - - @pytest.fixture() def client(monkeypatch): Base.metadata.create_all(bind=engine) diff --git a/backend/tests/test_branch_clause.py b/backend/tests/test_branch_clause.py index cf84fdf..6e52fd5 100644 --- a/backend/tests/test_branch_clause.py +++ b/backend/tests/test_branch_clause.py @@ -21,15 +21,6 @@ nodes may appear on C. python -m pytest tests/test_branch_clause.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient diff --git a/backend/tests/test_branch_forking.py b/backend/tests/test_branch_forking.py index 41cac9b..9fb1181 100644 --- a/backend/tests/test_branch_forking.py +++ b/backend/tests/test_branch_forking.py @@ -12,15 +12,6 @@ that makes borrowing possible lives in `lineage`. python -m pytest tests/test_branch_forking.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient @@ -29,9 +20,10 @@ from app import auth, limits, models from app.context import cursors, lineage from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts from app.routers import adventures +from fakes import ScriptedProvider + # `hp` moves freely. `mana` has a cooldown of 2 turns, so an incorrect # advance shows up as a change the referee should have rejected. SCHEMA = { @@ -50,22 +42,6 @@ modifier(text); """ -class ScriptedProvider: - last_usage = None - replies: list = [] - calls = 0 - prompts: list = [] - - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1) - ScriptedProvider.calls += 1 - ScriptedProvider.prompts.append((parts.system, parts.story)) - yield ("text", ScriptedProvider.replies[index]) - - @pytest.fixture() def client(monkeypatch): Base.metadata.create_all(bind=engine) diff --git a/backend/tests/test_branch_management.py b/backend/tests/test_branch_management.py index 27023ea..5c1f0b0 100644 --- a/backend/tests/test_branch_management.py +++ b/backend/tests/test_branch_management.py @@ -19,15 +19,6 @@ Two rules carry most of this file: python -m pytest tests/test_branch_management.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient @@ -37,22 +28,9 @@ from app import auth, limits, models, schemas from app.context import cursors, lineage from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts from app.routers import adventures - -class ScriptedProvider: - last_usage = None - replies: list = [] - calls = 0 - - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1) - ScriptedProvider.calls += 1 - yield ("text", ScriptedProvider.replies[index]) +from fakes import ScriptedProvider @pytest.fixture() diff --git a/backend/tests/test_bundle_v2.py b/backend/tests/test_bundle_v2.py index cb2ff28..e3c0682 100644 --- a/backend/tests/test_bundle_v2.py +++ b/backend/tests/test_bundle_v2.py @@ -23,15 +23,6 @@ importer a file that does disagree with itself. python -m pytest tests/test_bundle_v2.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient @@ -40,9 +31,10 @@ from app import auth, bundle, limits, models from app.context import lineage from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts from app.routers import adventures +from fakes import ScriptedProvider + SCHEMA = {"player": {"hp": {"min": 0, "max": 100, "initial": 100}}} # Ten gold a turn, so the stored gold total tells how many turns the @@ -58,20 +50,6 @@ modifier(text); OPENING = "You enter a cave." -class ScriptedProvider: - last_usage = None - replies: list = [] - calls = 0 - - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1) - ScriptedProvider.calls += 1 - yield ("text", ScriptedProvider.replies[index]) - - @pytest.fixture() def client(monkeypatch): Base.metadata.create_all(bind=engine) diff --git a/backend/tests/test_chat.py b/backend/tests/test_chat.py index df52166..42acfce 100644 --- a/backend/tests/test_chat.py +++ b/backend/tests/test_chat.py @@ -6,15 +6,6 @@ page. python -m pytest tests/test_chat.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient diff --git a/backend/tests/test_egress.py b/backend/tests/test_egress.py index 64cf5b5..372f72e 100644 --- a/backend/tests/test_egress.py +++ b/backend/tests/test_egress.py @@ -17,15 +17,6 @@ Two kinds of guard live here, and both are needed: python -m pytest tests/test_egress.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import json import random diff --git a/backend/tests/test_embedding_blob.py b/backend/tests/test_embedding_blob.py index c45a2c0..6fa82a5 100644 --- a/backend/tests/test_embedding_blob.py +++ b/backend/tests/test_embedding_blob.py @@ -7,15 +7,8 @@ for bit, and that the migration reaches every row, however many there are. python -m pytest tests/test_embedding_blob.py -v """ -import os import struct -import tempfile -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import json import random diff --git a/backend/tests/test_embedding_model_switch.py b/backend/tests/test_embedding_model_switch.py index 4e3db33..292d62d 100644 --- a/backend/tests/test_embedding_model_switch.py +++ b/backend/tests/test_embedding_model_switch.py @@ -19,15 +19,6 @@ garbage instead. python -m pytest tests/test_embedding_model_switch.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import asyncio import pytest diff --git a/backend/tests/test_guest_cleanup.py b/backend/tests/test_guest_cleanup.py index 995ba57..5f180d4 100644 --- a/backend/tests/test_guest_cleanup.py +++ b/backend/tests/test_guest_cleanup.py @@ -6,15 +6,8 @@ deleted. python -m pytest tests/test_guest_cleanup.py -v """ -import os -import tempfile from datetime import timedelta -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest from sqlalchemy import create_engine, event diff --git a/backend/tests/test_history_window.py b/backend/tests/test_history_window.py index d0a5635..54972c0 100644 --- a/backend/tests/test_history_window.py +++ b/backend/tests/test_history_window.py @@ -16,15 +16,6 @@ Two things must hold, and both are easy to break by accident: python -m pytest tests/test_history_window.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from sqlalchemy import event diff --git a/backend/tests/test_length_hint.py b/backend/tests/test_length_hint.py index 7477e29..0c26659 100644 --- a/backend/tests/test_length_hint.py +++ b/backend/tests/test_length_hint.py @@ -16,15 +16,8 @@ Two things are easy to break here: python -m pytest tests/test_length_hint.py -v """ -import os import re -import tempfile -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest diff --git a/backend/tests/test_memory_nodes.py b/backend/tests/test_memory_nodes.py index e2fa158..2d6f24f 100644 --- a/backend/tests/test_memory_nodes.py +++ b/backend/tests/test_memory_nodes.py @@ -19,15 +19,6 @@ exactly as `test_branch_clause.py` builds it. python -m pytest tests/test_memory_nodes.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import asyncio import pytest diff --git a/backend/tests/test_memory_retrieval.py b/backend/tests/test_memory_retrieval.py index 110902d..98980f2 100644 --- a/backend/tests/test_memory_retrieval.py +++ b/backend/tests/test_memory_retrieval.py @@ -14,15 +14,6 @@ failure that no error message would report: python -m pytest tests/test_memory_retrieval.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import asyncio from datetime import timedelta diff --git a/backend/tests/test_memory_settling.py b/backend/tests/test_memory_settling.py index fa0fdb8..f339b7e 100644 --- a/backend/tests/test_memory_settling.py +++ b/backend/tests/test_memory_settling.py @@ -23,14 +23,7 @@ happened. python -m pytest tests/test_memory_settling.py -v """ import asyncio -import os -import tempfile -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest diff --git a/backend/tests/test_netguard.py b/backend/tests/test_netguard.py index a0465f4..4b4e962 100644 --- a/backend/tests/test_netguard.py +++ b/backend/tests/test_netguard.py @@ -2,15 +2,6 @@ python -m pytest tests/test_netguard.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from app import auth, netguard diff --git a/backend/tests/test_prompt_caching.py b/backend/tests/test_prompt_caching.py index a84c0dd..88d51af 100644 --- a/backend/tests/test_prompt_caching.py +++ b/backend/tests/test_prompt_caching.py @@ -25,13 +25,7 @@ than assumed. python -m pytest tests/test_prompt_caching.py -v """ import os -import tempfile -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest diff --git a/backend/tests/test_ratelimit_hardening.py b/backend/tests/test_ratelimit_hardening.py index 60328e0..a83455f 100644 --- a/backend/tests/test_ratelimit_hardening.py +++ b/backend/tests/test_ratelimit_hardening.py @@ -9,15 +9,6 @@ also has an email-keyed throttle that no IP trick can weaken. python -m pytest tests/test_ratelimit_hardening.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from app import auth, limits diff --git a/backend/tests/test_retry_variants.py b/backend/tests/test_retry_variants.py index 4623bbf..fad75c6 100644 --- a/backend/tests/test_retry_variants.py +++ b/backend/tests/test_retry_variants.py @@ -4,15 +4,6 @@ switchable, restoring the world/script state that attempt produced. python -m pytest tests/test_retry_variants.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient @@ -20,9 +11,11 @@ from fastapi.testclient import TestClient from app import auth, limits, models from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts, ProviderError +from app.providers import ProviderError from app.routers import adventures +from fakes import ScriptedProvider + SCHEMA = {"player": {"hp": {"min": 0, "max": 100, "initial": 100}}} # Each turn spends 10 gold, so a double-applied or un-rolled-back attempt shows. @@ -35,26 +28,6 @@ modifier(text); """ -class ScriptedProvider: - """Streams the next canned reply each call, so successive retries differ.""" - last_usage = None - replies: list = [] - calls = 0 - prompts: list = [] # every assembled (system, story) pair, for context assertions - - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1) - ScriptedProvider.calls += 1 - ScriptedProvider.prompts.append((parts.system, parts.story)) - reply = ScriptedProvider.replies[index] - if isinstance(reply, Exception): - raise reply - yield ("text", reply) - - @pytest.fixture() def client(monkeypatch): Base.metadata.create_all(bind=engine) diff --git a/backend/tests/test_scenario_art.py b/backend/tests/test_scenario_art.py index 44d8180..3d7e46c 100644 --- a/backend/tests/test_scenario_art.py +++ b/backend/tests/test_scenario_art.py @@ -9,14 +9,7 @@ player's last line. python -m pytest tests/test_scenario_art.py -v """ import base64 -import os -import tempfile -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest from fastapi import Depends diff --git a/backend/tests/test_scenario_refresh.py b/backend/tests/test_scenario_refresh.py index d5c0d73..9e37532 100644 --- a/backend/tests/test_scenario_refresh.py +++ b/backend/tests/test_scenario_refresh.py @@ -3,15 +3,6 @@ story cards and stat schema back down over a running adventure's copy. python -m pytest tests/test_scenario_refresh.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient diff --git a/backend/tests/test_seed_sweep.py b/backend/tests/test_seed_sweep.py index c5affd0..0ff965d 100644 --- a/backend/tests/test_seed_sweep.py +++ b/backend/tests/test_seed_sweep.py @@ -8,14 +8,7 @@ deletes seeded rows no file claims any more. python -m pytest tests/test_seed_sweep.py -v """ import json -import os -import tempfile -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest diff --git a/backend/tests/test_snapshot_compression.py b/backend/tests/test_snapshot_compression.py index 20e9fd9..26b8e23 100644 --- a/backend/tests/test_snapshot_compression.py +++ b/backend/tests/test_snapshot_compression.py @@ -17,16 +17,9 @@ Three things must hold, and only the first is obvious: python -m pytest tests/test_snapshot_compression.py -v """ import json -import os import random -import tempfile import zlib -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest from sqlalchemy import text diff --git a/backend/tests/test_starter_adventure.py b/backend/tests/test_starter_adventure.py index 38e478a..3ca393e 100644 --- a/backend/tests/test_starter_adventure.py +++ b/backend/tests/test_starter_adventure.py @@ -8,14 +8,7 @@ empty account would. python -m pytest tests/test_starter_adventure.py -v """ import json -import os -import tempfile -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest from sqlalchemy import create_engine, event diff --git a/backend/tests/test_state_revert.py b/backend/tests/test_state_revert.py index 9fe18aa..998241b 100644 --- a/backend/tests/test_state_revert.py +++ b/backend/tests/test_state_revert.py @@ -11,17 +11,10 @@ only in their outcome. Run from the backend dir: python -m pytest tests/test_state_revert.py -v """ -import os -import tempfile # Point the app at a throwaway SQLite file before importing anything that # binds the engine at import time. `app.database` reads `AIDND_DB_PATH` # on import. -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest from fastapi import HTTPException diff --git a/backend/tests/test_story_tree_baseline.py b/backend/tests/test_story_tree_baseline.py index 36b7697..f5cf8b5 100644 --- a/backend/tests/test_story_tree_baseline.py +++ b/backend/tests/test_story_tree_baseline.py @@ -17,15 +17,6 @@ those subphases, the change is wrong, not the test. SP4 is the first subphase allowed to move it, and only for the variant-count semantics called out in plan/14. """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient @@ -33,9 +24,10 @@ from fastapi.testclient import TestClient from app import auth, limits, models from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts from app.routers import adventures +from fakes import ScriptedProvider + # A world-state schema, so the RPG layer is exercised rather than skipped. SCHEMA = {"player": {"hp": {"min": 0, "max": 100, "initial": 100}}} @@ -52,27 +44,6 @@ modifier(text); OPENING = "You enter a cave." -class ScriptedProvider: - """Streams the next canned reply each call, so successive turns differ.""" - last_usage = None - - replies: list = [] - calls = 0 - prompts: list = [] # every assembled (system, story) pair - - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1) - ScriptedProvider.calls += 1 - ScriptedProvider.prompts.append((parts.system, parts.story)) - reply = ScriptedProvider.replies[index] - if isinstance(reply, Exception): - raise reply - yield ("text", reply) - - def _make_world(monkeypatch, *, seeded_actions: int = 0): """Create a user, a scenario, and an adventure, with `seeded_actions` extra story actions written straight to the database. Paging tests need diff --git a/backend/tests/test_take_edit.py b/backend/tests/test_take_edit.py index 5300490..d02c433 100644 --- a/backend/tests/test_take_edit.py +++ b/backend/tests/test_take_edit.py @@ -20,15 +20,6 @@ be silently broken. python -m pytest tests/test_take_edit.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient @@ -36,22 +27,9 @@ from fastapi.testclient import TestClient from app import auth, limits, models from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts from app.routers import adventures - -class ScriptedProvider: - last_usage = None - replies: list = [] - calls = 0 - - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1) - ScriptedProvider.calls += 1 - yield ("text", ScriptedProvider.replies[index]) +from fakes import ScriptedProvider @pytest.fixture() diff --git a/backend/tests/test_take_parentage.py b/backend/tests/test_take_parentage.py index 47fbe84..0ce72e0 100644 --- a/backend/tests/test_take_parentage.py +++ b/backend/tests/test_take_parentage.py @@ -20,14 +20,7 @@ they mean. python -m pytest tests/test_take_parentage.py -v """ import json -import os -import tempfile -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest from fastapi import Depends @@ -36,22 +29,9 @@ from fastapi.testclient import TestClient from app import attempts, auth, limits, models from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts from app.routers import adventures - -class ScriptedProvider: - last_usage = None - replies: list = [] - calls = 0 - - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1) - ScriptedProvider.calls += 1 - yield ("text", ScriptedProvider.replies[index]) +from fakes import ScriptedProvider @pytest.fixture() diff --git a/backend/tests/test_take_state.py b/backend/tests/test_take_state.py index c02e1cb..49429d5 100644 --- a/backend/tests/test_take_state.py +++ b/backend/tests/test_take_state.py @@ -15,15 +15,6 @@ the line it was on. python -m pytest tests/test_take_state.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient @@ -31,9 +22,10 @@ from fastapi.testclient import TestClient from app import auth, limits, models from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts from app.routers import adventures +from fakes import ScriptedProvider + SCHEMA = {"player": {"hp": {"min": 0, "max": 100, "initial": 100}}} # Ten gold a turn, every turn. A number that only ever increases makes a @@ -48,20 +40,6 @@ modifier(text); """ -class ScriptedProvider: - last_usage = None - replies: list = [] - calls = 0 - - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - index = min(ScriptedProvider.calls, len(ScriptedProvider.replies) - 1) - ScriptedProvider.calls += 1 - yield ("text", ScriptedProvider.replies[index]) - - @pytest.fixture() def client(monkeypatch): Base.metadata.create_all(bind=engine) diff --git a/backend/tests/test_tree_migration.py b/backend/tests/test_tree_migration.py index e775b59..05746b8 100644 --- a/backend/tests/test_tree_migration.py +++ b/backend/tests/test_tree_migration.py @@ -16,14 +16,7 @@ place, would silently skip the DDL and test only half the change. python -m pytest tests/test_tree_migration.py -v """ import json -import os -import tempfile -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) import pytest from fastapi import Depends diff --git a/backend/tests/test_turn_flow_integration.py b/backend/tests/test_turn_flow_integration.py index c5745f4..0630065 100644 --- a/backend/tests/test_turn_flow_integration.py +++ b/backend/tests/test_turn_flow_integration.py @@ -7,15 +7,6 @@ adventure's stored gold total stays correct across play, undo, and retry. python -m pytest tests/test_turn_flow_integration.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient @@ -23,9 +14,10 @@ from fastapi.testclient import TestClient from app import auth, limits, models from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts from app.routers import adventures +from fakes import ScriptedProvider + GOLD_SCRIPT = """ const modifier = (text) => { state.gold = (state.gold || 0) + 10; @@ -35,14 +27,7 @@ modifier(text); """ -class FakeProvider: - """Stand-in for OpenAICompatibleProvider: streams one fixed line, no network.""" - last_usage = None - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - yield ("text", "The torch flickers as you press onward.") +AI_REPLY = "The torch flickers as you press onward." @pytest.fixture() @@ -66,7 +51,8 @@ def client(monkeypatch): setup.close() # Force a real, non-demo turn that uses the fake provider. - monkeypatch.setattr(adventures, "OpenAICompatibleProvider", FakeProvider) + ScriptedProvider.replies = [AI_REPLY] + monkeypatch.setattr(adventures, "OpenAICompatibleProvider", ScriptedProvider) monkeypatch.setattr(auth, "resolve_provider_config", lambda s: auth.ProviderConfig( "http://fake", "k", "test-model", False)) monkeypatch.setattr(limits, "rate_limit", lambda *a, **k: None) diff --git a/backend/tests/test_worldstate_integration.py b/backend/tests/test_worldstate_integration.py index 87a195f..1054344 100644 --- a/backend/tests/test_worldstate_integration.py +++ b/backend/tests/test_worldstate_integration.py @@ -4,15 +4,6 @@ undo rolling the world state back. python -m pytest tests/test_worldstate_integration.py -v """ -import os -import tempfile - -_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False) -_tmp.close() -os.environ["AIDND_DB_PATH"] = _tmp.name -os.environ.pop("AIDND_DATABASE_URL", None) -os.environ.pop("DATABASE_URL", None) - import pytest from fastapi import Depends from fastapi.testclient import TestClient @@ -20,9 +11,10 @@ from fastapi.testclient import TestClient from app import auth, limits, models from app.database import Base, SessionLocal, engine, get_db from app.main import app -from app.providers import PromptParts from app.routers import adventures +from fakes import ScriptedProvider + SCHEMA = { "player": {"hp": {"min": 0, "max": 100, "initial": 100, "max_delta_per_turn": 30}}, "npcs": { @@ -45,15 +37,6 @@ AI_REPLY = ( ) -class FakeProvider: - last_usage = None - def __init__(self, *a, **k): - pass - - async def generate(self, parts: PromptParts, *, temperature, max_tokens): - yield ("text", AI_REPLY) - - @pytest.fixture() def client(monkeypatch): Base.metadata.create_all(bind=engine) @@ -78,7 +61,8 @@ def client(monkeypatch): adv_id, user_id = adv.id, user.id setup.close() - monkeypatch.setattr(adventures, "OpenAICompatibleProvider", FakeProvider) + ScriptedProvider.replies = [AI_REPLY] + monkeypatch.setattr(adventures, "OpenAICompatibleProvider", ScriptedProvider) monkeypatch.setattr(auth, "resolve_provider_config", lambda s: auth.ProviderConfig( "http://fake", "k", "test-model", False)) monkeypatch.setattr(limits, "rate_limit", lambda *a, **k: None)