Skip to content
Open
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
12 changes: 12 additions & 0 deletions app/lib/backend/schema/gen/wrapped_task_integrations_wire.g.dart
Original file line number Diff line number Diff line change
Expand Up @@ -107,28 +107,40 @@ class GeneratedOAuthUrlResponse {
}

class GeneratedCreateTaskResponse {
final bool? ambiguous;
final String? error;
final String? errorCode;
final String? externalTaskId;
final bool? retryable;
final bool success;

const GeneratedCreateTaskResponse({
this.ambiguous,
this.error,
this.errorCode,
this.externalTaskId,
this.retryable,
required this.success,
});

factory GeneratedCreateTaskResponse.fromJson(Map<String, dynamic> json) {
return GeneratedCreateTaskResponse(
ambiguous: _readFieldValue<bool>(_readField(json, const ["ambiguous"]), "ambiguous", _readBool, requiredField: false, nullable: true),
error: _readFieldValue<String>(_readField(json, const ["error"]), "error", _readString, requiredField: false, nullable: true),
errorCode: _readFieldValue<String>(_readField(json, const ["error_code"]), "error_code", _readString, requiredField: false, nullable: true),
externalTaskId: _readFieldValue<String>(_readField(json, const ["external_task_id"]), "external_task_id", _readString, requiredField: false, nullable: true),
retryable: _readFieldValue<bool>(_readField(json, const ["retryable"]), "retryable", _readBool, requiredField: false, nullable: true),
success: _required(_readFieldValue<bool>(_readField(json, const ["success"]), "success", _readBool, requiredField: true, nullable: false), "success"),
);
}

Map<String, dynamic> toJson() {
return {
'ambiguous': ambiguous,
'error': error,
'error_code': errorCode,
'external_task_id': externalTaskId,
'retryable': retryable,
'success': success,
};
}
Expand Down
111 changes: 64 additions & 47 deletions backend/database/memories.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import copy

Check warning on line 1 in backend/database/memories.py

View workflow job for this annotation

GitHub Actions / Hygiene

Large changed file

backend/database/memories.py is 1299 lines; consider splitting files over 800 lines.
import hashlib
import json
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any, Callable, Dict, List, Optional, TypedDict, cast

Expand Down Expand Up @@ -34,6 +35,7 @@

memories_collection = 'memories'
users_collection = 'users'
_DELETE_BATCH_SIZE = 499


class MemoryDoc(TypedDict, total=False):
Expand Down Expand Up @@ -78,6 +80,15 @@
to_sha256: Optional[str]


@dataclass(frozen=True)
class LegacyMemoryDeleteResult:
memory_ids: List[str]

@property
def committed_count(self) -> int:
return len(self.memory_ids)


# Signature expected by ``prepare_for_read`` for the post-read decrypt hook. The
# concrete helper accepts/returns Optional[Dict] for direct call sites that may
# pass ``None``; at decorator sites we cast to this narrower contract.
Expand Down Expand Up @@ -562,24 +573,8 @@
return merge_evidence_sets(existing, incoming)


def delete_memories(uid: str, *, firestore_client: Any = None) -> None:
database = _get_db(firestore_client)
user_ref = database.collection(users_collection).document(uid)
memories_ref = user_ref.collection(memories_collection)
# Chunk deletes to stay under the Firestore 500-writes-per-batch limit. A user with more than
# 500 memories would otherwise make the single batch.commit() raise and delete nothing. Mirrors
# the chunking in unlock_all_memories.
batch = database.batch()
count = 0
for doc in memories_ref.stream():
batch.delete(doc.reference)
count += 1
if count >= 499: # Firestore batch limit is 500
batch.commit()
batch = database.batch()
count = 0
if count > 0:
batch.commit()
def delete_memories(uid: str, *, firestore_client: Any = None) -> LegacyMemoryDeleteResult:
return delete_all_memories(uid, firestore_client=firestore_client)


@prepare_for_read(decrypt_func=cast(_DecryptFunc, _prepare_memory_for_read))
Expand Down Expand Up @@ -882,57 +877,79 @@
)


def delete_memory(uid: str, memory_id: str, *, firestore_client: Any = None) -> None:
def delete_memory(uid: str, memory_id: str, *, firestore_client: Any = None) -> LegacyMemoryDeleteResult:
database = _get_db(firestore_client)
user_ref = database.collection(users_collection).document(uid)
memories_ref = user_ref.collection(memories_collection)
memory_ref = memories_ref.document(memory_id)
memory_ref.delete()
return _delete_memory_references(
[(memory_id, memory_ref)],
database=database,
)


def delete_memories_batch(uid: str, memory_ids: List[str], *, firestore_client: Any = None) -> None:
def delete_memories_batch(
uid: str,
memory_ids: List[str],
*,
firestore_client: Any = None,
) -> LegacyMemoryDeleteResult:
"""Delete multiple memories in a single batched Firestore write.

The router caps a batch-delete request at MEMORIES_BATCH_MAX (100), well under
Firestore's 500-writes-per-batch limit, but this helper mirrors the chunking in
delete_all_memories so it stays correct if it is ever reused for larger sets.
"""
if not memory_ids:
return
return LegacyMemoryDeleteResult(memory_ids=[])
database = _get_db(firestore_client)
user_ref = database.collection(users_collection).document(uid)
memories_ref = user_ref.collection(memories_collection)
batch = database.batch()
count = 0
for memory_id in memory_ids:
batch.delete(memories_ref.document(memory_id))
count += 1
if count >= 499: # Firestore batch limit is 500
batch.commit()
batch = database.batch()
count = 0
if count > 0:
batch.commit()
references = [(memory_id, memories_ref.document(memory_id)) for memory_id in dict.fromkeys(memory_ids)]
return _delete_memory_references(
references,
database=database,
)


def delete_all_memories(uid: str, *, firestore_client: Any = None) -> None:
def delete_all_memories(
uid: str,
*,
memory_ids: Optional[List[str]] = None,
firestore_client: Any = None,
) -> LegacyMemoryDeleteResult:
"""Delete one authoritative snapshot and return the exact committed IDs."""
database = _get_db(firestore_client)
user_ref = database.collection(users_collection).document(uid)
memories_ref = user_ref.collection(memories_collection)
# Chunk deletes to stay under the Firestore 500-writes-per-batch limit. Account deletion and
# "delete all memories" hit this for any user with more than 500 memories: the single
# batch.commit() would raise and remove nothing. Mirrors the chunking in unlock_all_memories.
batch = database.batch()
count = 0
for doc in memories_ref.stream():
batch.delete(doc.reference)
count += 1
if count >= 499: # Firestore batch limit is 500
batch.commit()
batch = database.batch()
count = 0
if count > 0:
references = (
[(memory_id, memories_ref.document(memory_id)) for memory_id in dict.fromkeys(memory_ids)]
if memory_ids is not None
else [(doc.id, doc.reference) for doc in memories_ref.stream()]
)
return _delete_memory_references(
references,
database=database,
)


def _delete_memory_references(
references: List[tuple[str, Any]],
*,
database: Any,
) -> LegacyMemoryDeleteResult:
if not references:
return LegacyMemoryDeleteResult(memory_ids=[])

committed_ids: List[str] = []
for offset in range(0, len(references), _DELETE_BATCH_SIZE):
chunk = references[offset : offset + _DELETE_BATCH_SIZE]
batch = database.batch()
for _memory_id, reference in chunk:
batch.delete(reference)
batch.commit()
committed_ids.extend(memory_id for memory_id, _reference in chunk)
return LegacyMemoryDeleteResult(memory_ids=committed_ids)


def ripple_source_deletion(uid: str, source_id: str, *, firestore_client: Any = None) -> Dict[str, Any]:
Expand Down
9 changes: 8 additions & 1 deletion backend/database/projection_repair.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,7 @@ def enqueue_projection_repairs(
collection_ref: Any = database.collection(users_collection).document(uid).collection(projection_repairs_collection)
repair_ids: List[str] = []
reasons_by_fact = _reasons_by_fact(mutations)
pending_writes = 0
for fact_id in fact_ids:
reasons = reasons_by_fact.get(fact_id, ['unknown'])
repair_id = f"{commit.get('commit_id')}:{fact_id}"
Expand All @@ -82,7 +83,13 @@ def enqueue_projection_repairs(
'updated_at': now,
},
)
batch.commit()
pending_writes += 1
if pending_writes >= 499:
batch.commit()
batch = database.batch()
pending_writes = 0
if pending_writes:
batch.commit()
return repair_ids


Expand Down
Loading
Loading