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"]