From 5371dfc91130103be63f2c3f95d1236dfd0598d7 Mon Sep 17 00:00:00 2001 From: Phil Merrell Date: Sun, 27 Sep 2026 11:54:04 -0600 Subject: [PATCH] fix: roll back KB sync change gates when the S3 stage fails The Drive and web re-crawl sync paths write sourceEtag/contentHash alongside previousChunkCount and stagedContentHash BEFORE staging to S3 (those two must precede the ObjectCreated event). When the put raised, the run failed but the gates stayed advanced, so every later run short-circuited as "unchanged" and the change was never ingested, on both legacy and managed knowledge bases. A failed stage now conditionally restores sourceEtag, contentHash, stagedContentHash and lastSyncedAt to their prior values, only if they still hold what this run wrote, so a newer run's write and the managed consumer's over-cap REMOVE are left alone. The crawler reports a new "stage_failed" outcome instead of letting the exception escape its worker task, marks a never-staged new page failed, and still walks the fetched page's links so its children aren't counted as misses. Co-Authored-By: Claude Opus 5.5 --- backend/src/apis/app_api/kb_sync/records.py | 63 +++++ backend/src/apis/app_api/kb_sync/worker.py | 49 +++- .../src/apis/app_api/web_sources/crawler.py | 44 ++- .../apis/app_api/web_sources/test_crawler.py | 38 +++ backend/tests/lambdas/test_kb_sync_worker.py | 252 ++++++++++++++++++ 5 files changed, 435 insertions(+), 11 deletions(-) diff --git a/backend/src/apis/app_api/kb_sync/records.py b/backend/src/apis/app_api/kb_sync/records.py index 64a865dfc..946ad8019 100644 --- a/backend/src/apis/app_api/kb_sync/records.py +++ b/backend/src/apis/app_api/kb_sync/records.py @@ -125,6 +125,69 @@ def update_document_sync_fields( ) +def rollback_document_sync_fields( + assistant_id: str, + document_id: str, + *, + written: Dict[str, str], + previous: Dict[str, Optional[str]], +) -> bool: + """Undo a pre-stage sync write whose S3 stage then failed. + + ``written`` maps each attribute to the value the sync just wrote; + ``previous`` maps it to the value it held before (``None`` = absent). + The change-detection gates (``sourceEtag``, ``contentHash``) have to be + written before staging only because they travel with the fields that + must precede the ObjectCreated event, so a stage that never happened + must not leave them advanced — the next run would match them and never + stage the change. ``stagedContentHash`` goes back too: left behind, it + names a version the object does not hold. + + Conditioned on every attribute still holding what was written, so a + newer write (another run, the consumer clearing the gates) is left + alone. Returns False when that condition refused the rollback. + """ + from botocore.exceptions import ClientError + + set_parts = [] + remove_parts = [] + conditions = [] + names: Dict[str, str] = {} + values: Dict[str, Any] = {} + for i, (attribute, value) in enumerate(written.items()): + names[f"#a{i}"] = attribute + values[f":w{i}"] = value + conditions.append(f"#a{i} = :w{i}") + prior = previous.get(attribute) + if prior is None: + remove_parts.append(f"#a{i}") + else: + set_parts.append(f"#a{i} = :p{i}") + values[f":p{i}"] = prior + if not conditions: + return True + + expression = [] + if set_parts: + expression.append("SET " + ", ".join(set_parts)) + if remove_parts: + expression.append("REMOVE " + ", ".join(remove_parts)) + try: + _table().update_item( + Key={"PK": f"AST#{assistant_id}", "SK": f"DOC#{document_id}"}, + UpdateExpression=" ".join(expression), + ConditionExpression=" AND ".join(conditions), + ExpressionAttributeNames=names, + ExpressionAttributeValues=values, + ) + except ClientError as exc: + if exc.response.get("Error", {}).get("Code") != "ConditionalCheckFailedException": + raise + logger.info(f"Document {document_id}'s sync fields changed since the failed stage; not rolling back") + return False + return True + + def clear_document_sync_policy_id(assistant_id: str, document_id: str) -> None: """Remove the SyncPolicy back-pointer when its policy is deleted.""" _table().update_item( diff --git a/backend/src/apis/app_api/kb_sync/worker.py b/backend/src/apis/app_api/kb_sync/worker.py index a615e3650..3d13c8089 100644 --- a/backend/src/apis/app_api/kb_sync/worker.py +++ b/backend/src/apis/app_api/kb_sync/worker.py @@ -212,16 +212,36 @@ async def _sync_drive_file(policy: SyncPolicy) -> Dict[str, Any]: # KB's consumer uses to tell this overwrite from a redelivery — BEFORE # staging, then overwrite the S3 object. previous_chunk_count = int(document.get("chunkCount") or 0) + synced_at = _now_timestamp() records.update_document_sync_fields( assistant_id, policy.source_ref, source_etag=new_etag, content_hash=content_hash, previous_chunk_count=previous_chunk_count, - last_synced_at=_now_timestamp(), + last_synced_at=synced_at, staged_content_hash=content_hash, ) - _stage_to_s3(document["s3Key"], downloaded.content, downloaded.content_type) + try: + _stage_to_s3(document["s3Key"], downloaded.content, downloaded.content_type) + except Exception: + # The gates advanced above but the bytes never reached S3; left + # as they are, the next run would match them and never stage + # this change. previousChunkCount is harmless to leave — the next + # changed run rewrites it before its own stage. + written = { + "sourceEtag": new_etag, + "contentHash": content_hash, + "stagedContentHash": content_hash, + "lastSyncedAt": synced_at, + } + records.rollback_document_sync_fields( + assistant_id, + policy.source_ref, + written=written, + previous={attribute: document.get(attribute) for attribute in written}, + ) + raise logger.info( f"Sync policy {policy.policy_id}: staged {len(downloaded.content)} changed bytes for " f"document {policy.source_ref} (prev chunks: {previous_chunk_count})" @@ -298,6 +318,9 @@ async def _sync_web_crawl(policy: SyncPolicy) -> Dict[str, Any]: web_docs[url] = item now = _now_timestamp() + # document_id -> (values written before staging, values they replaced), + # so a page whose stage then fails can be rolled back. + pre_stage_writes: Dict[str, Any] = {} async def on_result(url: str, document_id: str, outcome: str, etag, content_hash) -> None: if outcome == "changed": @@ -313,6 +336,13 @@ async def on_result(url: str, document_id: str, outcome: str, etag, content_hash last_synced_at=now, staged_content_hash=content_hash, ) + written = {"contentHash": content_hash, "stagedContentHash": content_hash, "lastSyncedAt": now} + if etag is not None: + written["sourceEtag"] = etag + pre_stage_writes[document_id] = ( + written, + {attribute: web_docs[url].get(attribute) for attribute in written}, + ) elif outcome == "unchanged": records.update_document_sync_fields( assistant_id, document_id, source_etag=etag, content_hash=content_hash, last_synced_at=now @@ -323,6 +353,13 @@ async def on_result(url: str, document_id: str, outcome: str, etag, content_hash records.update_document_sync_fields( assistant_id, document_id, content_hash=content_hash, last_synced_at=now ) + pre_stage_writes[document_id] = ({"contentHash": content_hash, "lastSyncedAt": now}, {}) + elif outcome == "stage_failed" and document_id in pre_stage_writes: + # The page's gates advanced above but its bytes never reached S3; + # roll them back so the next re-crawl stages it instead of + # hash-matching a version the knowledge base never received. + written, previous = pre_stage_writes.pop(document_id) + records.rollback_document_sync_fields(assistant_id, document_id, written=written, previous=previous) refresh = crawler.RefreshState( docs={ @@ -402,9 +439,13 @@ async def on_result(url: str, document_id: str, outcome: str, etag, content_hash logger.info( f"Sync policy {policy.policy_id}: re-crawl done — {refresh.changed} changed, " - f"{refresh.created} new, {refresh.unchanged} unchanged, {deleted} deleted" + f"{refresh.created} new, {refresh.unchanged} unchanged, {deleted} deleted, " + f"{refresh.stage_failed} failed to stage" ) - result = "changed" if (refresh.changed or refresh.created or deleted) else "unchanged" + # A page that failed to stage was counted as changed/created when it was + # emitted, but nothing reached the knowledge base. + staged = refresh.changed + refresh.created - refresh.stage_failed + result = "changed" if (staged or deleted) else "unchanged" return await _finish(policy, result) diff --git a/backend/src/apis/app_api/web_sources/crawler.py b/backend/src/apis/app_api/web_sources/crawler.py index 5b284f065..537d40ecb 100644 --- a/backend/src/apis/app_api/web_sources/crawler.py +++ b/backend/src/apis/app_api/web_sources/crawler.py @@ -152,7 +152,10 @@ class RefreshState: "unchanged" — 304 or identical content hash; nothing re-staged "changed" — existing doc, new bytes; emitted BEFORE staging so the worker can stash the previous chunk count first - "created" — page new to this crawl; emitted after staging + "created" — page new to this crawl; also emitted before staging + "stage_failed" — the S3 stage after a "changed"/"created" raised; the + worker rolls back the gate values it wrote, or the next + refresh would hash-match bytes that never reached S3 `seen_urls` collects every URL that survived the robots gate — the worker diffs it against `docs` for miss counting. Fetch failures ARE seen (a flaky page is not a missing page); robots-disallowed pages are @@ -167,6 +170,7 @@ class RefreshState: changed: int = 0 unchanged: int = 0 created: int = 0 + stage_failed: int = 0 async def _emit( self, @@ -182,6 +186,8 @@ async def _emit( self.unchanged += 1 elif outcome == "created": self.created += 1 + elif outcome == "stage_failed": + self.stage_failed += 1 if self.on_result is not None: await self.on_result(url, document_id, outcome, etag, content_hash) @@ -648,12 +654,36 @@ async def worker(url: str, depth: int) -> None: url, len(markdown.encode("utf-8")), ) - s3_key = await _put_markdown( - assistant_id=assistant_id, - document_id=document_id, - markdown=markdown, - filename=filename, - ) + try: + s3_key = await _put_markdown( + assistant_id=assistant_id, + document_id=document_id, + markdown=markdown, + filename=filename, + ) + except Exception as stage_err: + logger.warning("Staging %s to S3 failed: %s", url, stage_err) + if refresh is not None: + await refresh._emit(url, document_id, "stage_failed", etag, content_hash) + if existing is None: + await update_document_status( + assistant_id=assistant_id, + document_id=document_id, + status="failed", + error_message="The page could not be stored.", + error_details=str(stage_err)[:500], + ) + await increment_counters( + assistant_id=assistant_id, + crawl_id=crawl_id, + failed_delta=1, + ) + # The fetch itself succeeded, so its links are good. + # Not walking them would leave every page below this + # one unseen, and a refresh counts unseen pages as + # misses toward deletion. + await enqueue_links(html, url, depth) + return await update_document_import_metadata( assistant_id=assistant_id, document_id=document_id, diff --git a/backend/tests/apis/app_api/web_sources/test_crawler.py b/backend/tests/apis/app_api/web_sources/test_crawler.py index 3501c45dd..fd8177e91 100644 --- a/backend/tests/apis/app_api/web_sources/test_crawler.py +++ b/backend/tests/apis/app_api/web_sources/test_crawler.py @@ -650,3 +650,41 @@ async def test_refresh_robots_disallow_is_not_seen(recorder: _Recorder): # Robots said stop indexing: the URL is deliberately NOT seen, so the # worker's miss counter starts ticking toward removal. assert state.seen_urls == set() + + +@pytest.mark.asyncio +async def test_refresh_stage_failure_emits_stage_failed( + recorder: _Recorder, monkeypatch: pytest.MonkeyPatch +): + async def failing_put(*, assistant_id, document_id, markdown, filename): + raise RuntimeError("S3 PutObject failed") + + monkeypatch.setattr(crawler, "_put_markdown", failing_put) + pages = { + ROOT: '

changed words

n', + "https://example.com/new": "

new page

", + } + state, log = _refresh_state( + recorder, {ROOT: crawler.RefreshDoc(document_id="DOC-root", content_hash="different")} + ) + + await _run_refresh(recorder, pages, state) + + # Each pre-stage emit is followed by stage_failed, so the worker can roll + # back the gate values it wrote for bytes that never reached S3. + outcomes = sorted((outcome, url) for outcome, url, _ in log.events) + assert outcomes == [ + ("changed", ROOT), + ("created", "https://example.com/new"), + ("stage_failed", ROOT), + ("stage_failed", "https://example.com/new"), + ] + assert state.stage_failed == 2 + assert recorder.failed_delta == 2 + assert recorder.metadata_updates == [] + # The existing page keeps serving its last-good version; the new page, + # which has nothing staged at all, is marked failed. + new_doc_id = recorder.created_docs[0][0] + assert (new_doc_id, "failed") in recorder.status_updates + assert ("DOC-root", "failed") not in recorder.status_updates + assert recorder.finalized_status == "complete" diff --git a/backend/tests/lambdas/test_kb_sync_worker.py b/backend/tests/lambdas/test_kb_sync_worker.py index bdd6f0453..6086ba139 100644 --- a/backend/tests/lambdas/test_kb_sync_worker.py +++ b/backend/tests/lambdas/test_kb_sync_worker.py @@ -202,6 +202,120 @@ async def test_trashed_is_grace_skip(self, assistants_table, staged, token_ok, p assert adapter.download_calls == 0 +class TestStageFailure: + """A failed S3 stage must not leave the change-detection gates advanced — + otherwise the next run matches them and the change is never staged.""" + + @pytest.fixture() + def flaky_stage(self, monkeypatch): + """First stage raises (S3 put failure), later ones succeed.""" + calls = [] + + def stage(*args): + calls.append(args) + if len(calls) == 1: + raise RuntimeError("S3 PutObject failed") + + monkeypatch.setattr(worker, "_stage_to_s3", stage) + return calls + + def _doc(self, assistants_table, assistant_id): + return assistants_table.get_item(Key={"PK": f"AST#{assistant_id}", "SK": "DOC#doc-1"})["Item"] + + async def test_failed_stage_is_restaged_next_run( + self, assistants_table, flaky_stage, token_ok, provider_ok, monkeypatch + ): + adapter = FakeDriveAdapter(metadata={"version": "42", "trashed": False}, content=b"new bytes") + _use_adapter(monkeypatch, adapter) + assistant_id, _, policy = await _setup(assistants_table, etag="41", chunk_count=7) + + first = await worker.run_sync(_payload(assistant_id, policy)) + + assert first["result"] == "failed" + item = self._doc(assistants_table, assistant_id) + assert item["sourceEtag"] == "41" + assert "contentHash" not in item + assert "stagedContentHash" not in item + assert "lastSyncedAt" not in item + + # Same Drive version, same bytes: both gates must still miss. + second = await worker.run_sync(_payload(assistant_id, policy)) + + assert second["result"] == "changed" + assert len(flaky_stage) == 2 + assert flaky_stage[1][1] == b"new bytes" + item = self._doc(assistants_table, assistant_id) + assert item["sourceEtag"] == "42" + assert item["contentHash"] == worker._sha256(b"new bytes") + assert item["stagedContentHash"] == worker._sha256(b"new bytes") + + async def test_failed_stage_restores_previous_sync_values( + self, assistants_table, flaky_stage, token_ok, provider_ok, monkeypatch + ): + adapter = FakeDriveAdapter(metadata={"version": "42", "trashed": False}, content=b"new bytes") + _use_adapter(monkeypatch, adapter) + assistant_id, _, policy = await _setup(assistants_table, etag="41", content_hash="old-hash") + assistants_table.update_item( + Key={"PK": f"AST#{assistant_id}", "SK": "DOC#doc-1"}, + UpdateExpression="SET stagedContentHash = :h, lastSyncedAt = :t", + ExpressionAttributeValues={":h": "old-hash", ":t": "2026-01-01T00:00:00Z"}, + ) + + await worker.run_sync(_payload(assistant_id, policy)) + + item = self._doc(assistants_table, assistant_id) + assert item["sourceEtag"] == "41" + assert item["contentHash"] == "old-hash" + assert item["stagedContentHash"] == "old-hash" + assert item["lastSyncedAt"] == "2026-01-01T00:00:00Z" + + async def test_rollback_leaves_a_newer_write_alone(self, assistants_table): + from apis.app_api.kb_sync import records + + assistant_id, _, _ = await _setup(assistants_table, etag="41") + written = {"sourceEtag": "42", "contentHash": "h2", "stagedContentHash": "h2"} + records.update_document_sync_fields( + assistant_id, "doc-1", source_etag="42", content_hash="h2", staged_content_hash="h2" + ) + # A later run staged a newer version before this rollback landed. + records.update_document_sync_fields( + assistant_id, "doc-1", source_etag="43", content_hash="h3", staged_content_hash="h3" + ) + + rolled_back = records.rollback_document_sync_fields( + assistant_id, "doc-1", written=written, previous={"sourceEtag": "41"} + ) + + assert rolled_back is False + item = self._doc(assistants_table, assistant_id) + assert (item["sourceEtag"], item["contentHash"], item["stagedContentHash"]) == ("43", "h3", "h3") + + async def test_rollback_does_not_resurrect_gates_the_consumer_cleared(self, assistants_table): + """The managed consumer's over-cap refusal REMOVEs the gates to force a + re-stage; a late rollback must not write them back.""" + from apis.app_api.kb_sync import records + + assistant_id, _, _ = await _setup(assistants_table, etag="41", content_hash="h1") + records.update_document_sync_fields( + assistant_id, "doc-1", source_etag="42", content_hash="h2", staged_content_hash="h2" + ) + assistants_table.update_item( + Key={"PK": f"AST#{assistant_id}", "SK": "DOC#doc-1"}, + UpdateExpression="REMOVE sourceEtag, contentHash", + ) + + rolled_back = records.rollback_document_sync_fields( + assistant_id, + "doc-1", + written={"sourceEtag": "42", "contentHash": "h2", "stagedContentHash": "h2"}, + previous={"sourceEtag": "41", "contentHash": "h1"}, + ) + + assert rolled_back is False + item = self._doc(assistants_table, assistant_id) + assert "sourceEtag" not in item and "contentHash" not in item + + class TestFailureModes: async def test_not_found_strikes_counter(self, assistants_table, staged, token_ok, provider_ok, monkeypatch): adapter = FakeDriveAdapter(metadata_error=FileSourceNotFoundError("gone or unshared")) @@ -571,3 +685,141 @@ async def test_missing_root_doc_is_recreated(self, assistants_table, fake_crawl) item = assistants_table.get_item(Key={"PK": f"AST#{assistant_id}", "SK": f"DOC#{root_doc_id}"})["Item"] assert item["sourceFileId"] == self.ROOT assert item["sourceConnectorId"] == "web" + + async def test_failed_stage_rolls_back_page_gates(self, assistants_table, fake_crawl): + assistant_id, job, policy = await self._setup_crawl(assistants_table) + assistants_table.update_item( + Key={"PK": f"AST#{assistant_id}", "SK": "DOC#doc-web-0"}, + UpdateExpression="SET sourceEtag = :e, contentHash = :h", + ExpressionAttributeValues={":e": '"e1"', ":h": "hash1"}, + ) + fake_crawl["seen"] = [self.ROOT] + fake_crawl["emit"] = [ + (self.ROOT, "changed", '"e2"', "hash2"), + (self.ROOT, "stage_failed", '"e2"', "hash2"), + ] + + result = await worker.run_sync(self._payload(assistant_id, policy, job)) + + # Nothing reached the knowledge base, so the run changed nothing. + assert result["result"] == "unchanged" + item = assistants_table.get_item(Key={"PK": f"AST#{assistant_id}", "SK": "DOC#doc-web-0"})["Item"] + assert item["sourceEtag"] == '"e1"' + assert item["contentHash"] == "hash1" + assert "stagedContentHash" not in item + + # The next re-crawl still sees the old gates, so it will re-stage. + fake_crawl["emit"] = [] + await worker.run_sync(self._payload(assistant_id, policy, job)) + refresh_doc = fake_crawl["captured"]["refresh"].docs[self.ROOT] + assert (refresh_doc.source_etag, refresh_doc.content_hash) == ('"e1"', "hash1") + + +class TestWebCrawlStageFailureEndToEnd: + """Worker + the REAL crawler (httpx on a MockTransport, S3 put stubbed): + a page whose stage fails is staged again by the next re-crawl, for both + a changed page and a page new to the crawl.""" + + ROOT = "https://example.com/" + NEW = "https://example.com/new" + + @pytest.fixture() + def site(self, monkeypatch): + import httpx + + from apis.app_api.web_sources import crawler + + pages = { + self.ROOT: '

fresh root words

n', + self.NEW: "

a brand new page

", + } + + def handler(request): + url = str(request.url) + if url in pages: + return httpx.Response(200, text=pages[url], headers={"content-type": "text/html"}) + return httpx.Response(404, text="") + + monkeypatch.setattr( + crawler, "_default_client", lambda: httpx.AsyncClient(transport=httpx.MockTransport(handler)) + ) + puts = {"fail": True, "ok": []} + + async def put_markdown(*, assistant_id, document_id, markdown, filename): + if puts["fail"]: + raise RuntimeError("S3 PutObject failed") + puts["ok"].append(document_id) + return f"assistants/{assistant_id}/documents/{document_id}/{filename}" + + monkeypatch.setattr(crawler, "_put_markdown", put_markdown) + return puts + + async def test_failed_stages_are_restaged_next_run(self, assistants_table, site): + from apis.app_api.web_sources.crawl_repository import create_crawl_job, finalize_crawl + from apis.app_api.web_sources.models import CrawlSettings + + assistant = await create_assistant( + owner_id=USER_ID, owner_name="U", name="A", description="d", + instructions="i", vector_index_id="assistants-index", + ) + assistant_id = assistant.assistant_id + job = await create_crawl_job( + assistant_id=assistant_id, root_url=self.ROOT, + settings=CrawlSettings(max_depth=1, max_pages=10, min_delay_seconds=0, max_delay_seconds=0), + started_by_user_id=USER_ID, + ) + await finalize_crawl(assistant_id=assistant_id, crawl_id=job.crawl_id, status="complete") + await create_document( + assistant_id=assistant_id, filename="root.md", content_type="text/markdown", + size_bytes=1, s3_key=f"assistants/{assistant_id}/documents/doc-root/root.md", + document_id="doc-root", + provenance=DocumentProvenance( + source_connector_id="web", source_adapter_key="http", + source_file_id=self.ROOT, imported_by_user_id=USER_ID, + ), + ) + assistants_table.update_item( + Key={"PK": f"AST#{assistant_id}", "SK": "DOC#doc-root"}, + UpdateExpression="SET contentHash = :h", + ExpressionAttributeValues={":h": "old-root-hash"}, + ) + policy = await create_sync_policy( + assistant_id=assistant_id, source_type="web_crawl", source_ref=job.crawl_id, + interval="daily", created_by_user_id=USER_ID, + ) + payload = { + "policyId": policy.policy_id, "assistantId": assistant_id, + "sourceType": "web_crawl", "sourceRef": job.crawl_id, + } + + first = await worker.run_sync(payload) + + assert first["result"] == "unchanged" + assert site["ok"] == [] + root = assistants_table.get_item(Key={"PK": f"AST#{assistant_id}", "SK": "DOC#doc-root"})["Item"] + assert root["contentHash"] == "old-root-hash" + assert "stagedContentHash" not in root + new_docs = [ + item for item in _document_items(assistants_table, assistant_id) + if item.get("sourceFileId") == self.NEW + ] + assert len(new_docs) == 1 + assert "contentHash" not in new_docs[0] + assert new_docs[0]["status"] == "failed" + + site["fail"] = False + second = await worker.run_sync(payload) + + assert second["result"] == "changed" + assert sorted(site["ok"]) == sorted(["doc-root", new_docs[0]["documentId"]]) + root = assistants_table.get_item(Key={"PK": f"AST#{assistant_id}", "SK": "DOC#doc-root"})["Item"] + assert root["contentHash"] != "old-root-hash" + assert root["stagedContentHash"] == root["contentHash"] + + +def _document_items(assistants_table, assistant_id): + from boto3.dynamodb.conditions import Key + + return assistants_table.query( + KeyConditionExpression=Key("PK").eq(f"AST#{assistant_id}") & Key("SK").begins_with("DOC#") + )["Items"]