diff --git a/e2e/test_admin_vscode_chat.py b/e2e/test_admin_vscode_chat.py index e0d5f1d54..c2a28a0ad 100644 --- a/e2e/test_admin_vscode_chat.py +++ b/e2e/test_admin_vscode_chat.py @@ -53,6 +53,7 @@ def startup(route): response = route.fetch() payload = response.json() payload["startup"]["integrations"]["vscode-chat"] = dict(progress) + payload["startup"]["messaging"]["state"] = "starting" route.fulfill(response=response, json=payload) def status(route): @@ -74,9 +75,10 @@ def retry(route): page.route("**/admin/api/integrations/vscode-chat/refresh", retry) page.goto(f"{admin_base_url}/admin/integrations") # Startup reconciliation can temporarily disable Retry during the click. - # Exercise the action after initial polling has settled. + # Wait for this integration without waiting for unrelated startup work. page.wait_for_function( - "state.startup && !state.startupRequest && state.startupTimer === null" + "state.startup?.startup?.integrations?.['vscode-chat']?.state === 'failed'" + " && !vscodeChatIntegration.busy" ) main = page.locator("#openVSCodeChatIntegration") secondary = page.locator("#retryVSCodeChatIntegration") diff --git a/src/free_claude_code/runtime/code_sessions_sqlite.py b/src/free_claude_code/runtime/code_sessions_sqlite.py index c9846f327..0681609c5 100644 --- a/src/free_claude_code/runtime/code_sessions_sqlite.py +++ b/src/free_claude_code/runtime/code_sessions_sqlite.py @@ -4,7 +4,6 @@ import json import sqlite3 from collections.abc import Callable, Sequence -from contextlib import closing from free_claude_code.application.code_sessions.models import ( ACTIVE_RUN_STATUSES, @@ -376,7 +375,6 @@ def _write_prompt(connection: sqlite3.Connection, prompt: CodePrompt) -> None: class SQLiteCodeStore: def __init__(self, database: SQLiteDatabase) -> None: self.database = database - self._path = database.path self._started = False self._lifecycle = asyncio.Lock() @@ -386,7 +384,7 @@ async def start(self) -> None: return try: await self.database.start() - await self.database.work(lambda: self._execute(self._recover)) + await self._execute(self._recover, write=True) except sqlite3.IntegrityError as exc: raise CodeConflictError( "Saved Code history conflicts with its schema." @@ -401,13 +399,11 @@ async def close(self) -> None: async with self._lifecycle: self._started = False - def _execute[T](self, operation: Callable[[sqlite3.Connection], T]) -> T: + async def _execute[T]( + self, operation: Callable[[sqlite3.Connection], T], *, write: bool + ) -> T: try: - with closing(sqlite3.connect(self._path, timeout=10)) as connection: - connection.row_factory = sqlite3.Row - connection.execute("PRAGMA foreign_keys = ON") - with connection: - return operation(connection) + return await self.database.run(operation, write=write) except sqlite3.IntegrityError as exc: raise CodeConflictError( "This Code operation conflicts with existing session state." @@ -415,10 +411,12 @@ def _execute[T](self, operation: Callable[[sqlite3.Connection], T]) -> T: except sqlite3.Error as exc: raise CodeUnavailableError("Code session storage is unavailable.") from exc - async def _run[T](self, operation: Callable[[sqlite3.Connection], T]) -> T: + async def _run[T]( + self, operation: Callable[[sqlite3.Connection], T], *, write: bool + ) -> T: if not self._started: raise CodeUnavailableError("Code session storage is closed.") - return await self.database.work(lambda: self._execute(operation)) + return await self._execute(operation, write=write) def _recover(self, connection: sqlite3.Connection) -> None: connection.execute( @@ -435,7 +433,6 @@ def _recover(self, connection: sqlite3.Connection) -> None: async def create(self, session: CodeSession) -> CodeSession: def operation(connection: sqlite3.Connection) -> CodeSession: - connection.execute("BEGIN IMMEDIATE") if connection.execute( "SELECT 1 FROM code_deleted WHERE id = ?", (session.id,) ).fetchone(): @@ -455,10 +452,12 @@ def operation(connection: sqlite3.Connection) -> CodeSession: _insert(connection, "code_sessions", session) return session - return await self._run(operation) + return await self._run(operation, write=True) async def get_session(self, session_id: str) -> CodeSession: - return await self._run(lambda connection: _session(connection, session_id)) + return await self._run( + lambda connection: _session(connection, session_id), write=False + ) async def list_sessions( self, cursor: tuple[int, str] | None, limit: int, query: str = "" @@ -484,7 +483,7 @@ def operation(connection: sqlite3.Connection) -> CodePage: last = sessions[-1] if sessions and len(rows) > limit else None return CodePage(sessions, (last.updated_at, last.id) if last else None) - return await self._run(operation) + return await self._run(operation, write=False) async def pending_deletions(self) -> tuple[CodeSession, ...]: return await self._run( @@ -493,7 +492,8 @@ async def pending_deletions(self) -> tuple[CodeSession, ...]: for row in connection.execute( "SELECT * FROM code_sessions WHERE status != 'ready'" ) - ) + ), + write=False, ) async def is_deleted(self, session_id: str) -> bool: @@ -503,7 +503,8 @@ async def is_deleted(self, session_id: str) -> bool: "SELECT 1 FROM code_deleted WHERE id = ?", (session_id,) ).fetchone() is not None - ) + ), + write=False, ) async def get_run(self, session_id: str, run_id: str) -> CodeRun | None: @@ -514,7 +515,7 @@ def operation(connection: sqlite3.Connection) -> CodeRun | None: ).fetchone() return _record(CodeRun, row) if row else None - return await self._run(operation) + return await self._run(operation, write=False) async def runs(self, session_id: str) -> tuple[CodeRun, ...]: return await self._run( @@ -524,7 +525,8 @@ async def runs(self, session_id: str) -> tuple[CodeRun, ...]: "SELECT * FROM code_runs WHERE session_id = ? ORDER BY ordinal", (session_id,), ) - ) + ), + write=False, ) async def latest_run(self, session_id: str) -> CodeRun | None: @@ -535,7 +537,7 @@ def operation(connection: sqlite3.Connection) -> CodeRun | None: ).fetchone() return _record(CodeRun, row) if row else None - return await self._run(operation) + return await self._run(operation, write=False) async def items( self, session_id: str, before: tuple[int, int] | None, limit: int | None @@ -546,10 +548,9 @@ async def item_page( self, session_id: str, before: tuple[int, int] | None, limit: int | None ) -> CodeItemPage: def operation(connection: sqlite3.Connection) -> CodeItemPage: - connection.execute("BEGIN") return _item_page(connection, session_id, before, limit) - return await self._run(operation) + return await self._run(operation, write=False) async def read_history( self, @@ -558,7 +559,6 @@ async def read_history( include_item_ids: Sequence[str], ) -> CodeHistory: def operation(connection: sqlite3.Connection) -> CodeHistory: - connection.execute("BEGIN") session = _session(connection, session_id) run = _latest_run(connection, session_id) active = run is not None and run.status in ACTIVE_RUN_STATUSES @@ -609,11 +609,10 @@ def operation(connection: sqlite3.Connection) -> CodeHistory: page.next_before, ) - return await self._run(operation) + return await self._run(operation, write=False) async def execution_seed(self, session_id: str) -> CodeExecutionSeed: def operation(connection: sqlite3.Connection) -> CodeExecutionSeed: - connection.execute("BEGIN") session = _session(connection, session_id) run = _latest_run(connection, session_id) sequence = connection.execute( @@ -636,7 +635,7 @@ def operation(connection: sqlite3.Connection) -> CodeExecutionSeed: _by_ids(connection, CodeRun, session_id, tuple(ids)), ) - return await self._run(operation) + return await self._run(operation, write=False) async def get_native_item( self, session_id: str, turn_id: str, item_id: str @@ -648,7 +647,7 @@ def operation(connection: sqlite3.Connection) -> CodeItem | None: ).fetchone() return _record(CodeItem, row) if row else None - return await self._run(operation) + return await self._run(operation, write=False) async def run_items(self, session_id: str, run_id: str) -> tuple[CodeItem, ...]: return await self._run( @@ -658,14 +657,16 @@ async def run_items(self, session_id: str, run_id: str) -> tuple[CodeItem, ...]: "SELECT * FROM code_items WHERE session_id = ? AND run_id = ? ORDER BY sequence", (session_id, run_id), ) - ) + ), + write=False, ) async def get_prompt(self, session_id: str, prompt_id: str) -> CodePrompt | None: return await self._run( lambda connection: next( iter(_by_ids(connection, CodePrompt, session_id, (prompt_id,))), None - ) + ), + write=False, ) async def has_prompt_request( @@ -678,7 +679,8 @@ async def has_prompt_request( (session_id, generation, json.dumps(request_id)), ).fetchone() is not None - ) + ), + write=False, ) async def has_native_turn(self, session_id: str, turn_id: str) -> bool: @@ -689,7 +691,8 @@ async def has_native_turn(self, session_id: str, turn_id: str) -> bool: (session_id, turn_id), ).fetchone() is not None - ) + ), + write=False, ) async def prompts(self, session_id: str) -> tuple[CodePrompt, ...]: @@ -700,14 +703,14 @@ async def prompts(self, session_id: str) -> tuple[CodePrompt, ...]: "SELECT * FROM code_prompts WHERE session_id = ? ORDER BY id", (session_id,), ) - ) + ), + write=False, ) async def update_settings( self, session: CodeSession, expected_revision: int ) -> CodeSession: def operation(connection: sqlite3.Connection) -> CodeSession: - connection.execute("BEGIN IMMEDIATE") previous = _session(connection, session.id) if previous.status != "ready": raise CodeConflictError("This session is being deleted.") @@ -720,13 +723,12 @@ def operation(connection: sqlite3.Connection) -> CodeSession: _write_session(connection, session, expected_revision, settings=True) return _session(connection, session.id) - return await self._run(operation) + return await self._run(operation, write=True) async def admit_run( self, session: CodeSession, run: CodeRun, item: CodeItem, expected_revision: int ) -> tuple[CodeSession, CodeRun]: def operation(connection: sqlite3.Connection) -> tuple[CodeSession, CodeRun]: - connection.execute("BEGIN IMMEDIATE") previous = _session(connection, session.id) row = connection.execute( "SELECT * FROM code_runs WHERE session_id = ? AND id = ?", @@ -768,13 +770,12 @@ def operation(connection: sqlite3.Connection) -> tuple[CodeSession, CodeRun]: _insert(connection, "code_items", item) return _session(connection, session.id), admitted - return await self._run(operation) + return await self._run(operation, write=True) async def claim_prompt( self, session_id: str, prompt_id: str, response_id: str, generation: str ) -> CodePrompt: def operation(connection: sqlite3.Connection) -> CodePrompt: - connection.execute("BEGIN IMMEDIATE") row = connection.execute( "SELECT * FROM code_prompts WHERE session_id = ? AND id = ?", (session_id, prompt_id), @@ -794,7 +795,7 @@ def operation(connection: sqlite3.Connection) -> CodePrompt: _write_prompt(connection, claimed) return claimed - return await self._run(operation) + return await self._run(operation, write=True) async def save_progress( self, @@ -806,7 +807,6 @@ async def save_progress( prompts: Sequence[CodePrompt] = (), ) -> None: def operation(connection: sqlite3.Connection) -> None: - connection.execute("BEGIN IMMEDIATE") _write_session(connection, session, expected_revision, settings=False) if run is not None: if run.session_id != session.id: @@ -821,14 +821,13 @@ def operation(connection: sqlite3.Connection) -> None: raise CodeConflictError("This prompt belongs to another session.") _write_prompt(connection, prompt) - await self._run(operation) + await self._run(operation, write=True) async def delete(self, session_id: str) -> None: def operation(connection: sqlite3.Connection) -> None: - connection.execute("BEGIN IMMEDIATE") connection.execute( "INSERT OR IGNORE INTO code_deleted(id) VALUES (?)", (session_id,) ) connection.execute("DELETE FROM code_sessions WHERE id = ?", (session_id,)) - await self._run(operation) + await self._run(operation, write=True) diff --git a/src/free_claude_code/runtime/sqlite_database.py b/src/free_claude_code/runtime/sqlite_database.py index 2c74da316..647b5d8b8 100644 --- a/src/free_claude_code/runtime/sqlite_database.py +++ b/src/free_claude_code/runtime/sqlite_database.py @@ -1,7 +1,7 @@ -"""Offline relocation and transactional schema initialization of FCC's database. +"""Connection, transaction, and lifetime ownership of FCC's shared database. -The application holds the shared owner lock until both features close. Connections -opened here never escape initialization, including when migration fails. +The application holds the owner lock until both features close. Each connection +is opened, used, and closed within initialization or one synchronous transaction. """ import asyncio @@ -228,6 +228,12 @@ def _initialize(self) -> None: def execute[T]( self, operation: Callable[[sqlite3.Connection], T], *, write: bool = True ) -> T: + """Run one transaction inside an admitted worker, closing its connection. + + Callbacks consume the connection synchronously and return materialized + results. They must not begin or finish the outer transaction themselves. + Legacy import uses this primitive for several transactions in one worker. + """ with closing(_connect(self.path)) as connection: connection.execute("BEGIN IMMEDIATE" if write else "BEGIN") try: @@ -242,6 +248,7 @@ def execute[T]( async def run[T]( self, operation: Callable[[sqlite3.Connection], T], *, write: bool = True ) -> T: + """Admit and supervise one transaction through committed result delivery.""" return await self.work(lambda: self.execute(operation, write=write)) async def work[T](self, operation: Callable[[], T]) -> T: diff --git a/tests/runtime/test_code_sessions_sqlite.py b/tests/runtime/test_code_sessions_sqlite.py index db5ffb6b5..f0f899511 100644 --- a/tests/runtime/test_code_sessions_sqlite.py +++ b/tests/runtime/test_code_sessions_sqlite.py @@ -161,6 +161,159 @@ async def _session(store): ) +@pytest.mark.asyncio +async def test_closed_database_uses_code_storage_error(store): + await store.database.close() + with pytest.raises(CodeUnavailableError, match="storage is unavailable") as error: + await store.get_session("missing") + assert isinstance(error.value.__cause__, sqlite3.OperationalError) + + +@pytest.mark.asyncio +async def test_code_read_uses_explicit_connection_policy(store, monkeypatch): + session = await _session(store) + read_session = code_store_module._session + + def observe(connection, session_id): + assert connection.autocommit is True + assert connection.in_transaction + assert connection.execute("PRAGMA foreign_keys").fetchone()[0] == 1 + assert connection.execute("SELECT 1 AS value").fetchone()["value"] == 1 + return read_session(connection, session_id) + + monkeypatch.setattr(code_store_module, "_session", observe) + assert await store.get_session(session.id) == session + + +@pytest.mark.asyncio +async def test_code_reads_do_not_reserve_writer_while_wal_write_is_open(store): + session = await _session(store) + with closing(sqlite3.connect(store.database.path, autocommit=True)) as writer: + writer.execute("BEGIN IMMEDIATE") + try: + writer.execute( + "UPDATE code_sessions SET title = 'uncommitted' WHERE id = ?", + (session.id,), + ) + saved = await asyncio.wait_for(store.get_session(session.id), 3) + history = await asyncio.wait_for( + store.read_history(session.id, None, ()), 3 + ) + assert saved == history.session == session + finally: + writer.execute("ROLLBACK") + + +@pytest.mark.asyncio +async def test_recovery_rolls_back_runs_when_prompt_expiration_fails(store): + session, run = await _admit(store, await _session(store)) + prompt = CodePrompt( + id="prompt", + session_id=session.id, + generation="g", + request_id=1, + kind="question", + form={}, + raw={}, + ) + item = CodeItem( + id=prompt.id, + session_id=session.id, + run_id=run.id, + sequence=2, + kind="prompt", + complete=True, + ) + await store.save_progress( + session, session.revision, items=(item,), prompts=(prompt,) + ) + await store.close() + await store.database.run( + lambda connection: connection.execute( + "CREATE TRIGGER reject_expiration BEFORE UPDATE ON code_prompts " + "BEGIN SELECT RAISE(ABORT, 'recovery failure'); END" + ) + ) + try: + with pytest.raises(CodeConflictError) as error: + await store.start() + assert isinstance(error.value.__cause__, sqlite3.IntegrityError) + states = await store.database.run( + lambda connection: ( + connection.execute("SELECT status FROM code_runs").fetchone()[0], + connection.execute("SELECT status FROM code_prompts").fetchone()[0], + ), + write=False, + ) + assert states == (run.status, prompt.status) + finally: + await store.database.run( + lambda connection: connection.execute("DROP TRIGGER reject_expiration") + ) + await store.start() + assert (await store.get_run(session.id, run.id)).status == "interrupted" + assert (await store.get_prompt(session.id, prompt.id)).status == "expired" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fail", [False, True]) +async def test_cancelled_code_write_delivers_outcome_before_database_close( + store, database_factory, monkeypatch, tmp_path, fail +): + entered, release = threading.Event(), threading.Event() + closing_started = asyncio.Event() + insert = code_store_module._insert + session = CodeSession(id="new", cwd="/work", model="provider/model") + + def blocked(connection, table, record): + insert(connection, table, record) + entered.set() + assert release.wait(5) + if fail: + raise sqlite3.OperationalError("injected write failure") + + async def close_database(): + closing_started.set() + await store.database.close() + + monkeypatch.setattr(code_store_module, "_insert", blocked) + writing = asyncio.create_task(store.create(session)) + closing_task = None + try: + assert await asyncio.to_thread(entered.wait, 3) + writing.cancel() + closing_task = asyncio.create_task(close_database()) + await asyncio.wait_for(closing_started.wait(), 3) + assert not closing_task.done() + with pytest.raises(CodeUnavailableError): + await store.get_session(session.id) + other = database_factory(store.database.path, tmp_path / "code.lock") + with pytest.raises(sqlite3.OperationalError, match="another FCC"): + await other.start() + release.set() + if fail: + with pytest.raises(CodeUnavailableError) as error: + await writing + assert isinstance(error.value.__cause__, sqlite3.OperationalError) + else: + assert await writing == session + await asyncio.wait_for(closing_task, 3) + reopened = database_factory(store.database.path, tmp_path / "code.lock") + await reopened.start() + saved = await reopened.run( + lambda connection: connection.execute( + "SELECT id FROM code_sessions" + ).fetchall(), + write=False, + ) + assert [row["id"] for row in saved] == ([] if fail else [session.id]) + finally: + release.set() + await asyncio.gather( + writing, *([closing_task] if closing_task else []), return_exceptions=True + ) + + @pytest.mark.asyncio async def test_history_snapshot_decodes_only_selected_records_and_chunks_includes( store, monkeypatch