Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 10 additions & 1 deletion src/architectai/session/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
)
from architectai.orchestrator.state import SessionState
from architectai.schemas.plan import Plan
from architectai.schemas.research_findings import ResearchFindings
from architectai.session.store import RedisSessionStore

_logger = structlog.get_logger(__name__)
Expand Down Expand Up @@ -230,6 +231,8 @@ async def refine(self, session_id: UUID, refinement: str) -> None:
"Refinement is identical to your previous input. "
"Please describe what you'd like to change."
)
original_plan = state.plan
original_findings = state.findings
lock_key = f"session:{session_id}:refine_lock"
acquired = await self._store.acquire_lock(lock_key, REFINE_LOCK_TTL_MS)
if not acquired:
Expand All @@ -254,14 +257,18 @@ async def refine(self, session_id: UUID, refinement: str) -> None:
)
await self._store.set(new_state)

task = asyncio.create_task(self._run_refinement(session_id, new_state))
task = asyncio.create_task(
self._run_refinement(session_id, new_state, original_plan, original_findings)
)
_background_tasks.add(task)
task.add_done_callback(_background_tasks.discard)

async def _run_refinement(
self,
session_id: UUID,
new_state: SessionState,
original_plan: Plan | None,
original_findings: ResearchFindings | None,
) -> None:
lock_key = f"session:{session_id}:refine_lock"
try:
Expand Down Expand Up @@ -313,6 +320,8 @@ async def _run_refinement(
)
noop_state = new_state.model_copy(
update={
"plan": original_plan,
"findings": original_findings,
"status": "complete",
"progress_message": None,
"updated_at": self._clock(),
Expand Down
132 changes: 132 additions & 0 deletions tests/integration/test_refine_roundtrip.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,132 @@
"""Integration tests for SessionManager.refine() round-trip with real state persistence.

Closes issue #106.
"""

from __future__ import annotations

import asyncio
from uuid import UUID

import pytest
from fakeredis import FakeAsyncRedis

from architectai.agents.fixtures import DISCOVERY_FIXTURES
from architectai.session.manager import PlanOutcome, RefinementInProgressError, SessionManager
from architectai.session.store import RedisSessionStore
from tests.stubs import StubDiscoveryAgent, StubPlanningAgent, StubResearchAgent

_ORIGINAL_BRIEF = DISCOVERY_FIXTURES["web_saas_ats"].final_brief
_MODIFIED_BRIEF_SSO = _ORIGINAL_BRIEF.model_copy(
update={"product_one_liner": "niche ATS for small independent recruiting agencies — with SSO"}
)


def _make_manager_and_store(
discovery: StubDiscoveryAgent,
planning_fixture: str,
) -> tuple[SessionManager, RedisSessionStore, StubPlanningAgent]:
redis_client = FakeAsyncRedis(decode_responses=True)
store = RedisSessionStore(redis_client, ttl_seconds=3600)
research = StubResearchAgent("web_saas_ats_complete")
planning = StubPlanningAgent(planning_fixture)
manager = SessionManager(store, discovery, research, planning)
return manager, store, planning


async def _drive_to_complete(manager: SessionManager, raw_idea: str) -> UUID:
session_id, _ = await manager.create(raw_idea)
outcome = await manager.turn(session_id, "test answer")
while not isinstance(outcome, PlanOutcome):
outcome = await manager.turn(session_id, "test answer")
return session_id


@pytest.mark.integration
async def test_refine_full_pipeline_runs_and_persists_final_state():
discovery = StubDiscoveryAgent("web_saas_ats", re_derive_brief_result=_MODIFIED_BRIEF_SSO)
manager, store, _ = _make_manager_and_store(discovery, "web_saas_resolved")
session_id = await _drive_to_complete(manager, "i want to build an ATS for small firms")
await manager.refine(session_id, "add SSO support")
await asyncio.sleep(0)
await asyncio.sleep(0)
await asyncio.sleep(0)
await asyncio.sleep(0)
state = await store.get(session_id)
assert state is not None
assert state.status == "complete"
assert state.plan is not None
assert state.progress_message is None
assert state.iteration == 1
assert len(state.cost_tracker.records) > 0


@pytest.mark.integration
async def test_refine_noop_skips_pipeline_and_restores_complete():
discovery = StubDiscoveryAgent("web_saas_ats")
manager, store, _ = _make_manager_and_store(discovery, "web_saas_resolved")
session_id = await _drive_to_complete(manager, "i want to build an ATS")
original_state = await store.get(session_id)
assert original_state is not None
original_plan = original_state.plan
await manager.refine(session_id, "just a minor tweak")
await asyncio.sleep(0)
await asyncio.sleep(0)
await asyncio.sleep(0)
await asyncio.sleep(0)
state = await store.get(session_id)
assert state is not None
assert state.status == "complete"
assert state.progress_message is None
assert state.plan == original_plan


@pytest.mark.integration
async def test_refine_lock_is_released_after_completion():
discovery = StubDiscoveryAgent("web_saas_ats")
manager, store, _ = _make_manager_and_store(discovery, "web_saas_resolved")
session_id = await _drive_to_complete(manager, "i want to build an ATS")
await manager.refine(session_id, "add export feature")
await asyncio.sleep(0)
await asyncio.sleep(0)
await asyncio.sleep(0)
await asyncio.sleep(0)
lock_key = f"session:{session_id}:refine_lock"
assert await store.acquire_lock(lock_key, 1000) is True
await store.release_lock(lock_key)


@pytest.mark.integration
async def test_refine_rejects_at_iteration_cap():
discovery = StubDiscoveryAgent("web_saas_ats")
manager, store, _ = _make_manager_and_store(discovery, "web_saas_resolved")
session_id = await _drive_to_complete(manager, "i want to build an ATS")
state = await store.get(session_id)
assert state is not None
await store.set(state.model_copy(update={"iteration": 5}))
with pytest.raises(ValueError, match="Refinement limit reached"):
await manager.refine(session_id, "change something")


@pytest.mark.integration
async def test_refine_rejects_duplicate_of_last_history_answer():
discovery = StubDiscoveryAgent("web_saas_ats")
manager, store, _ = _make_manager_and_store(discovery, "web_saas_resolved")
session_id = await _drive_to_complete(manager, "i want to build an ATS")
state = await store.get(session_id)
assert state is not None
last_answer = state.history[-1].answer
with pytest.raises(ValueError, match="identical to your previous input"):
await manager.refine(session_id, last_answer)


@pytest.mark.integration
async def test_refine_raises_409_when_lock_already_held():
discovery = StubDiscoveryAgent("web_saas_ats")
manager, store, _ = _make_manager_and_store(discovery, "web_saas_resolved")
session_id = await _drive_to_complete(manager, "i want to build an ATS")
lock_key = f"session:{session_id}:refine_lock"
await store.acquire_lock(lock_key, 120_000)
with pytest.raises(RefinementInProgressError):
await manager.refine(session_id, "add caching layer")
await store.release_lock(lock_key)
Loading