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
63 changes: 63 additions & 0 deletions backend/src/apis/app_api/kb_sync/records.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
49 changes: 45 additions & 4 deletions backend/src/apis/app_api/kb_sync/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -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})"
Expand Down Expand Up @@ -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":
Expand All @@ -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
Expand All @@ -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={
Expand Down Expand Up @@ -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)


Expand Down
44 changes: 37 additions & 7 deletions backend/src/apis/app_api/web_sources/crawler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -167,6 +170,7 @@ class RefreshState:
changed: int = 0
unchanged: int = 0
created: int = 0
stage_failed: int = 0

async def _emit(
self,
Expand All @@ -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)

Expand Down Expand Up @@ -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,
Expand Down
38 changes: 38 additions & 0 deletions backend/tests/apis/app_api/web_sources/test_crawler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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: '<html><body><p>changed words</p><a href="/new">n</a></body></html>',
"https://example.com/new": "<html><body><p>new page</p></body></html>",
}
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"
Loading