From 4ea54d07e5d3e3e03fb3a393fe06e0c1abc1b143 Mon Sep 17 00:00:00 2001 From: ZachL111 Date: Tue, 28 Jul 2026 08:48:15 -0700 Subject: [PATCH 1/3] Harden boundary delivery and data convergence --- .../gen/wrapped_task_integrations_wire.g.dart | 12 + backend/database/memories.py | 111 ++-- backend/database/projection_repair.py | 9 +- backend/database/redis_db.py | 128 +++- backend/database/vector_db.py | 19 +- backend/database/webhook_health.py | 21 + backend/routers/memories.py | 117 +++- backend/routers/pusher.py | 330 ++++++++-- backend/routers/task_integrations.py | 6 + backend/routers/users.py | 12 +- backend/testing/e2e/test_task_integrations.py | 17 +- backend/testing/e2e/test_webhooks.py | 5 +- backend/testing/workflow_contracts.json | 65 +- .../tests/unit/test_async_app_integrations.py | 95 ++- .../unit/test_async_http_infrastructure.py | 29 +- backend/tests/unit/test_async_webhooks.py | 102 ++- .../tests/unit/test_dev_api_lock_bypass.py | 15 +- .../unit/test_developer_memory_adapter.py | 49 +- .../test_listen_finalization_cloud_tasks.py | 336 +++++++++- backend/tests/unit/test_lock_bypass_fixes.py | 24 +- .../tests/unit/test_mcp_search_memories.py | 20 +- backend/tests/unit/test_memories_batch.py | 68 +- .../tests/unit/test_memories_batch_delete.py | 29 +- .../unit/test_memories_delete_batch_chunk.py | 52 +- backend/tests/unit/test_memory_ledger.py | 57 +- .../tests/unit/test_memory_service_parity.py | 193 +++++- .../unit/test_pusher_ghost_connections.py | 31 +- backend/tests/unit/test_pusher_heartbeat.py | 14 +- .../tests/unit/test_pusher_readiness_drain.py | 4 +- .../unit/test_redis_db_cache_serialization.py | 85 ++- .../test_speaker_identification_delivery.py | 238 +++++++ ...st_task_integration_due_date_validation.py | 30 + .../tests/unit/test_task_integrations_ops.py | 137 +++++ .../test_tools_rest_memory_runtime_adapter.py | 19 - backend/tests/unit/test_tools_router.py | 86 ++- .../unit/test_users_webhook_url_validation.py | 37 +- .../tests/unit/test_webhook_auto_disable.py | 18 +- .../unit/test_x_memory_extraction_retry.py | 21 + .../unit/utils/test_listen_pusher_session.py | 479 ++++++++++++++- backend/utils/app_integrations.py | 83 ++- backend/utils/http_client.py | 18 +- backend/utils/listen_pusher_session.py | 579 +++++++++++++++--- backend/utils/memory/memory_service.py | 240 ++++++-- backend/utils/other/storage.py | 6 +- .../utils/retrieval/tool_services/memories.py | 45 +- backend/utils/speaker_identification.py | 63 +- backend/utils/task_integrations_ops.py | 175 ++++-- backend/utils/webhooks.py | 56 +- backend/utils/x_connector.py | 12 +- .../src/renderer/src/lib/omiApi.generated.ts | 3 + docs/api-reference/app-client-openapi.json | 33 + .../backend/listen_pusher_pipeline.mdx | 59 +- .../lib/services/omi-api/omiApi.generated.ts | 3 + web/app/src/lib/omiApi.generated.ts | 3 + .../src/lib/omiApi.generated.ts | 3 + 55 files changed, 3911 insertions(+), 590 deletions(-) create mode 100644 backend/tests/unit/test_speaker_identification_delivery.py diff --git a/app/lib/backend/schema/gen/wrapped_task_integrations_wire.g.dart b/app/lib/backend/schema/gen/wrapped_task_integrations_wire.g.dart index 52f22e88100..81249334421 100644 --- a/app/lib/backend/schema/gen/wrapped_task_integrations_wire.g.dart +++ b/app/lib/backend/schema/gen/wrapped_task_integrations_wire.g.dart @@ -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 json) { return GeneratedCreateTaskResponse( + ambiguous: _readFieldValue(_readField(json, const ["ambiguous"]), "ambiguous", _readBool, requiredField: false, nullable: true), error: _readFieldValue(_readField(json, const ["error"]), "error", _readString, requiredField: false, nullable: true), + errorCode: _readFieldValue(_readField(json, const ["error_code"]), "error_code", _readString, requiredField: false, nullable: true), externalTaskId: _readFieldValue(_readField(json, const ["external_task_id"]), "external_task_id", _readString, requiredField: false, nullable: true), + retryable: _readFieldValue(_readField(json, const ["retryable"]), "retryable", _readBool, requiredField: false, nullable: true), success: _required(_readFieldValue(_readField(json, const ["success"]), "success", _readBool, requiredField: true, nullable: false), "success"), ); } Map toJson() { return { + 'ambiguous': ambiguous, 'error': error, + 'error_code': errorCode, 'external_task_id': externalTaskId, + 'retryable': retryable, 'success': success, }; } diff --git a/backend/database/memories.py b/backend/database/memories.py index 88e41f54479..9ddb8c4b054 100644 --- a/backend/database/memories.py +++ b/backend/database/memories.py @@ -1,6 +1,7 @@ import copy import hashlib import json +from dataclasses import dataclass from datetime import datetime, timezone from typing import Any, Callable, Dict, List, Optional, TypedDict, cast @@ -29,6 +30,7 @@ class FirestoreNotFound(Exception): memories_collection = 'memories' users_collection = 'users' +_DELETE_BATCH_SIZE = 499 class MemoryDoc(TypedDict, total=False): @@ -73,6 +75,15 @@ class MemoryDoc(TypedDict, total=False): 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. @@ -403,24 +414,8 @@ def _merge_evidence( # type: ignore[reportUnusedFunction] # reserved: thin ali 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)) @@ -723,15 +718,23 @@ def write_projection(transaction: Any) -> None: ) -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 @@ -739,41 +742,55 @@ def delete_memories_batch(uid: str, memory_ids: List[str], *, firestore_client: 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]: diff --git a/backend/database/projection_repair.py b/backend/database/projection_repair.py index 9cbf6505e44..b8eeeabc4fc 100644 --- a/backend/database/projection_repair.py +++ b/backend/database/projection_repair.py @@ -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}" @@ -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 diff --git a/backend/database/redis_db.py b/backend/database/redis_db.py index a784607297c..2d0ad3c8013 100644 --- a/backend/database/redis_db.py +++ b/backend/database/redis_db.py @@ -1,8 +1,11 @@ import ast import base64 +import hashlib import json import os -from typing import Any, Callable, Dict, List, Optional, TypeVar, Union, cast +import secrets +import threading +from typing import Any, Callable, Dict, List, Optional, Tuple, TypeVar, Union, cast from datetime import datetime, timedelta, timezone import redis @@ -28,6 +31,24 @@ password=os.getenv('REDIS_DB_PASSWORD'), health_check_interval=30, ) +_pusher_delivery_r: Any = None +_pusher_delivery_r_lock = threading.Lock() + +_PUSHER_DELIVERY_COMPLETE_LUA = """ +if redis.call('GET', KEYS[1]) ~= ARGV[1] then + return 0 +end +redis.call('SET', KEYS[1], 'done', 'EX', ARGV[2]) +return 1 +""" + +_PUSHER_DELIVERY_ABANDON_LUA = """ +if redis.call('GET', KEYS[1]) ~= ARGV[1] then + return 0 +end +redis.call('DEL', KEYS[1]) +return 1 +""" T = TypeVar("T") @@ -841,7 +862,7 @@ def remove_conversation_summary_app_id(app_id: str) -> bool: # Lua script: atomic increment + TTL in a single round-trip. # Returns [current_count, ttl_remaining]. Sets TTL on first hit # and self-heals any key that lost its TTL (prevents permanent buckets). -_RATE_LIMIT_LUA = r.register_script(""" +_RATE_LIMIT_LUA_SOURCE = """ local key = KEYS[1] local window = tonumber(ARGV[1]) local current = redis.call('INCR', key) @@ -854,7 +875,8 @@ def remove_conversation_summary_app_id(app_id: str) -> bool: ttl = window end return {current, ttl} -""") +""" +_RATE_LIMIT_LUA = r.register_script(_RATE_LIMIT_LUA_SOURCE) def check_rate_limit(key: str, policy: str, max_requests: int, window: int) -> tuple[bool, int, int]: @@ -883,7 +905,7 @@ def check_rate_limit(key: str, policy: str, max_requests: int, window: int) -> t # Burst uses a sorted set keyed by timestamp-ms for sliding-window accuracy, # trimmed on every call (O(log n)). Daily char counter auto-expires at midnight # UTC (caller passes seconds_until_midnight_utc as the TTL). -_TTS_RATE_LIMIT_LUA = r.register_script(""" +_TTS_RATE_LIMIT_LUA_SOURCE = """ local burst_key = KEYS[1] local daily_key = KEYS[2] local now_ms = tonumber(ARGV[1]) @@ -911,7 +933,8 @@ def check_rate_limit(key: str, policy: str, max_requests: int, window: int) -> t redis.call('EXPIRE', daily_key, daily_ttl) end return {0, 0} -""") +""" +_TTS_RATE_LIMIT_LUA = r.register_script(_TTS_RATE_LIMIT_LUA_SOURCE) def _seconds_until_midnight_utc() -> int: @@ -958,6 +981,101 @@ def try_acquire_listen_lock(uid: str, ttl: int = 7) -> bool: return result is not None +def _get_pusher_delivery_redis() -> Any: + """Return a lazy, tightly bounded Redis client for realtime ACK fencing.""" + global _pusher_delivery_r + if _pusher_delivery_r is not None: + return _pusher_delivery_r + with _pusher_delivery_r_lock: + if _pusher_delivery_r is None: + _pusher_delivery_r = redis.Redis( + host=cast(str, _redis_host), + port=int(_redis_port_env) if _redis_port_env is not None else 6379, + username='default', + password=os.getenv('REDIS_DB_PASSWORD'), + health_check_interval=30, + socket_connect_timeout=0.5, + socket_timeout=0.5, + retry_on_timeout=False, + ) + return _pusher_delivery_r + + +def _pusher_delivery_key(uid: str, delivery_id: str) -> str: + digest = hashlib.sha256(delivery_id.encode('utf-8')).hexdigest() + return f'users:{uid}:pusher_delivery:{digest}' + + +def begin_pusher_delivery( + uid: str, + delivery_id: str, + lease_ttl: int = 600, + *, + redis_client: Any = None, +) -> Tuple[str, Optional[str]]: + """Acquire a short processing lease for a stable realtime delivery.""" + client = redis_client or _get_pusher_delivery_redis() + key = _pusher_delivery_key(uid, delivery_id) + lease_token = secrets.token_urlsafe(18) + processing_value = f'processing:{lease_token}' + try: + if client.set(key, processing_value, ex=lease_ttl, nx=True) is not None: + return 'claimed', lease_token + state = client.get(key) + if state is not None and _decode_redis_value(state) == 'done': + return 'done', None + return 'busy', None + except Exception as exc: + logger.warning('pusher delivery lease unavailable uid=%s error=%s', uid, type(exc).__name__) + return 'unavailable', None + + +def complete_pusher_delivery( + uid: str, + delivery_id: str, + lease_token: str, + retention_ttl: int = 604800, + *, + redis_client: Any = None, +) -> bool: + """Replace a processing lease with a bounded done marker.""" + client = redis_client or _get_pusher_delivery_redis() + try: + result = client.eval( + _PUSHER_DELIVERY_COMPLETE_LUA, + 1, + _pusher_delivery_key(uid, delivery_id), + f'processing:{lease_token}', + retention_ttl, + ) + return int(result) == 1 + except Exception as exc: + logger.warning('pusher delivery completion unavailable uid=%s error=%s', uid, type(exc).__name__) + return False + + +def abandon_pusher_delivery( + uid: str, + delivery_id: str, + lease_token: str, + *, + redis_client: Any = None, +) -> bool: + """Release this worker's lease after an effect fails or is cancelled.""" + client = redis_client or _get_pusher_delivery_redis() + try: + result = client.eval( + _PUSHER_DELIVERY_ABANDON_LUA, + 1, + _pusher_delivery_key(uid, delivery_id), + f'processing:{lease_token}', + ) + return int(result) == 1 + except Exception as exc: + logger.warning('pusher delivery abandon unavailable uid=%s error=%s', uid, type(exc).__name__) + return False + + def try_acquire_client_device_write_lock(uid: str, client_device_id: str, ttl: int = 600) -> bool: """Throttle client_devices registry upserts to once per (uid, device) every `ttl` seconds.""" try: diff --git a/backend/database/vector_db.py b/backend/database/vector_db.py index 4fb2ca0ac9e..858fed7a0cb 100644 --- a/backend/database/vector_db.py +++ b/backend/database/vector_db.py @@ -387,6 +387,14 @@ class VectorCandidateQueryResult: rejected_count: int = 0 +def _reported_upsert_count(result: Any) -> int | None: + if isinstance(result, dict): + count = result.get('upserted_count') + else: + count = getattr(result, 'upserted_count', None) + return count if type(count) is int and count >= 0 else None + + def upsert_memory_vector( uid: str, memory_id: str, @@ -426,6 +434,9 @@ def upsert_memory_vector( } res = index.upsert(vectors=[data], namespace=MEMORIES_NAMESPACE) logger.info(f'upsert_memory_vector {memory_id} {res}') + if _reported_upsert_count(res) == 0: + logger.warning('upsert_memory_vector returned zero writes memory_id=%s', memory_id) + return None return vector @@ -481,7 +492,8 @@ def upsert_memory_vectors_batch(uid: str, items: List[Dict[str, Any]]) -> int: ) res = index.upsert(vectors=payload, namespace=MEMORIES_NAMESPACE) logger.info(f'upsert_memory_vectors_batch count={len(payload)} {res}') - return len(payload) + reported_count = _reported_upsert_count(res) + return len(payload) if reported_count is None else reported_count def find_similar_memories( @@ -665,17 +677,18 @@ def query_memory_vector_candidates( return VectorCandidateQueryResult(hits=hits, rejected_count=rejected_count) -def delete_memory_vector(uid: str, memory_id: str) -> None: +def delete_memory_vector(uid: str, memory_id: str) -> bool: """ Delete a memory vector from Pinecone. """ if index is None: logger.warning('Pinecone index not initialized, skipping memory vector delete') - return + return False vector_id = f'{uid}-{memory_id}' result = index.delete(ids=[vector_id], namespace=MEMORIES_NAMESPACE) logger.info(f'delete_memory_vector {vector_id} {result}') + return True def enqueue_projection_repair(uid: str, fact_id: str, reason: str, source_commit_id: str | None = None) -> List[str]: diff --git a/backend/database/webhook_health.py b/backend/database/webhook_health.py index e865ad4392f..038317d5cdc 100644 --- a/backend/database/webhook_health.py +++ b/backend/database/webhook_health.py @@ -428,3 +428,24 @@ def record_dev_webhook_success(uid: str, wtype: object): r.expire(key, _HEALTH_TTL) except Exception as e: logger.warning(f'record_dev_webhook_success redis error uid={uid} type={wtype}: {e}') + + +def reset_dev_webhook_health(uid: str, wtype: object): + """Clear delivery failures without claiming that an HTTP request succeeded.""" + try: + wtype_str = getattr(wtype, 'value') if hasattr(wtype, 'value') else str(wtype) + key = f'dev_webhook_health:{uid}:{wtype_str}' + r.hset( + key, + mapping={ + 'failure_count': '0', + 'last_failure_at': '', + 'last_success_at': '', + 'last_status': '', + 'last_error': '', + 'disabled': '0', + }, + ) + r.expire(key, _HEALTH_TTL) + except Exception as e: + logger.warning(f'reset_dev_webhook_health redis error uid={uid} type={wtype}: {e}') diff --git a/backend/routers/memories.py b/backend/routers/memories.py index de489cfb06b..26e5d2e63a7 100644 --- a/backend/routers/memories.py +++ b/backend/routers/memories.py @@ -94,6 +94,18 @@ class ReviewResolutionResponse(BaseModel): _MEMORY_DEVICE_SCOPE_SUPPORTED_HEADER = 'X-Omi-Memory-Device-Scope-Supported' +def _single_vector_write_succeeded(result: Any) -> bool: + if result is None or result is False: + return False + if isinstance(result, (list, tuple)) and not result: + return False + return True + + +def _batch_vector_write_succeeded(result: Any, expected_count: int) -> bool: + return type(result) is int and result == expected_count + + @dataclass(frozen=True) class V3GetRuntime: """Lazy, overrideable F4 runtime bundle for GET `/v3/memories`. @@ -364,25 +376,30 @@ def _mirror_delete_into_legacy(uid: str, memory_ids: List[str], *, db_client: An def _purge_legacy_memories(uid: str) -> None: - """Delete every legacy memory for a user, plus its vectors. + """Delete one authoritative legacy snapshot and converge its vectors.""" - Collects ids before the Firestore delete so the Pinecone vectors can be purged too — - otherwise orphaned vectors become search noise nothing ever cleans up. - """ - memory_ids: List[str] = [] - offset = 0 - batch_size = 1000 - while True: - memories = memories_db.get_memories(uid, limit=batch_size, offset=offset, include_invalidated=True) - if not memories: - break - memory_ids.extend([memory_id for m in memories if isinstance((memory_id := m.get('id')), str) and memory_id]) - offset += batch_size - - memories_db.delete_all_memories(uid) - - if memory_ids: - delete_memory_vectors_batch(uid, memory_ids) + delete_result = memories_db.delete_all_memories(uid) + memory_ids = delete_result.memory_ids + if not memory_ids: + return + + try: + vector_deleted_count = delete_memory_vectors_batch(uid, memory_ids) + except Exception: + logger.exception( + "legacy delete-all vector purge failed uid=%s count=%d", + sanitize_pii(uid), + len(memory_ids), + ) + return + + if vector_deleted_count != len(memory_ids): + logger.warning( + "legacy delete-all vector purge was partial uid=%s expected=%d actual=%r", + sanitize_pii(uid), + len(memory_ids), + vector_deleted_count, + ) def _mirror_delete_all_into_legacy(uid: str, *, db_client: Any) -> None: @@ -482,7 +499,7 @@ async def create_memory( raise HTTPException(status_code=503, detail="Service temporarily unavailable") try: - await run_blocking( + projection_result = await run_blocking( postprocess_executor, upsert_memory_vector, uid, @@ -491,9 +508,14 @@ async def create_memory( memory_db.category.value, memory_db.subject_entity_id, ) + if not _single_vector_write_succeeded(projection_result): + logger.warning( + "Vector upsert returned no write uid=%s memory_id=%s (memory saved, vector missing)", + uid, + memory_db.id, + ) except Exception: logger.exception("Vector upsert failed uid=%s memory_id=%s (memory saved, vector missing)", uid, memory_db.id) - return _legacy_memory_response(memory_db) @@ -623,7 +645,7 @@ async def create_memories_batch( # single-create path) so a slow embeddings/Pinecone call can't starve the # FastAPI sync threadpool. try: - await run_blocking( + projection_result = await run_blocking( postprocess_executor, upsert_memory_vectors_batch, uid, @@ -637,11 +659,17 @@ async def create_memories_batch( for m in memory_dbs ], ) + if not _batch_vector_write_succeeded(projection_result, len(memory_dbs)): + logger.warning( + "Batch vector upsert returned partial/no write uid=%s expected=%s actual=%r", + uid, + len(memory_dbs), + projection_result, + ) except Exception: logger.exception( "Batch vector upsert failed uid=%s count=%s (memories saved, vectors missing)", uid, len(memory_dbs) ) - return _legacy_batch_memories_response(memory_dbs) @@ -888,9 +916,24 @@ def delete_memories_batch( if memory.get('is_locked', False): raise HTTPException(status_code=402, detail='A paid plan is required to access this memory.') - memories_db.delete_memories_batch(uid, memory_ids) + delete_result = memories_db.delete_memories_batch(uid, memory_ids) + if delete_result.committed_count != len(memory_ids): + logger.error( + "Firestore batch delete count mismatch uid=%s expected=%d actual=%r", + uid, + len(memory_ids), + delete_result.committed_count, + ) + raise HTTPException(status_code=503, detail="Service temporarily unavailable") try: - delete_memory_vectors_batch(uid, memory_ids) + vector_deleted_count = delete_memory_vectors_batch(uid, memory_ids) + if vector_deleted_count != len(memory_ids): + logger.warning( + "Vector batch delete returned partial/no write uid=%s expected=%d actual=%r", + uid, + len(memory_ids), + vector_deleted_count, + ) except Exception: logger.exception("Vector batch delete failed uid=%s count=%d (Firestore already deleted)", uid, len(memory_ids)) return {'status': 'ok'} @@ -913,9 +956,23 @@ def delete_memory( return {'status': 'ok'} _validate_memory(uid, memory_id) - memories_db.delete_memory(uid, memory_id) + delete_result = memories_db.delete_memory(uid, memory_id) + if delete_result.committed_count != 1: + logger.error( + "Firestore delete count mismatch uid=%s memory_id=%s actual=%r", + uid, + memory_id, + delete_result.committed_count, + ) + raise HTTPException(status_code=503, detail="Service temporarily unavailable") try: - delete_memory_vector(uid, memory_id) + projection_result = delete_memory_vector(uid, memory_id) + if projection_result is not True: + logger.warning( + "Vector delete returned no write uid=%s memory_id=%s (Firestore deleted)", + uid, + memory_id, + ) except Exception: logger.exception("Vector delete failed uid=%s memory_id=%s (Firestore deleted)", uid, memory_id) return {'status': 'ok'} @@ -984,13 +1041,19 @@ def edit_memory( # vector keeps matching the OLD text — a silent staleness bug that breaks the # "constantly updated brain" (search would still surface the pre-edit fact). try: - upsert_memory_vector( + projection_result = upsert_memory_vector( uid, memory_id, mutation_value, memory.get('category', 'system'), subject_entity_id=memory.get('subject_entity_id'), ) + if not _single_vector_write_succeeded(projection_result): + logger.warning( + "Vector upsert returned no write uid=%s memory_id=%s (memory edited, vector stale)", + uid, + memory_id, + ) except Exception: logger.exception("Vector upsert failed uid=%s memory_id=%s (memory edited, vector stale)", uid, memory_id) return {'status': 'ok'} diff --git a/backend/routers/pusher.py b/backend/routers/pusher.py index 83935b5f47c..5b92e6d21af 100644 --- a/backend/routers/pusher.py +++ b/backend/routers/pusher.py @@ -1,6 +1,7 @@ import struct import asyncio import json +import logging import time from collections import deque from typing import Any, Dict, List, Optional, TypedDict, cast @@ -11,6 +12,7 @@ import database.conversations as conversations_db from database import conversation_finalization_jobs as finalization_jobs_db +from database import redis_db from database import users as users_db from services.conversation_finalization import final_attempt_failed from utils.apps import is_audio_bytes_app_enabled @@ -42,8 +44,8 @@ from utils.metrics import PUSHER_ACTIVE_WS_CONNECTIONS from utils.readiness import ReadinessGate from utils.observability.journeys import JourneyAttempt, JourneyOutcome, record_capture_finalization_terminal +from utils.observability.fallback import record_fallback from utils.speaker_identification import extract_speaker_samples -import logging logger = logging.getLogger(__name__) @@ -80,6 +82,7 @@ # before being force-cancelled. Prevents hung GCS uploads or webhook calls from # blocking cleanup indefinitely. BG_DRAIN_TIMEOUT = 30.0 # seconds +PUSHER_DELIVERY_DRAIN_OPCODE = 107 def pusher_session_outcome(close_code: int, *, application_failed: bool = False) -> JourneyOutcome: @@ -96,11 +99,13 @@ class _SpeakerSampleRequest(TypedDict): conversation_id: str segment_ids: List[str] queued_at: float + delivery_id: Optional[str] class _TranscriptQueueItem(TypedDict): segments: List[Dict[str, Any]] memory_id: Optional[str] + delivery_id: Optional[str] class _AudioBytesQueueItem(TypedDict): @@ -124,6 +129,7 @@ async def _process_conversation_task( byok_keys: Optional[Dict[str, str]] = None, finalization_job_id: Optional[str] = None, dispatch_generation: Optional[int] = None, + send_lock: Optional[asyncio.Lock] = None, ) -> None: """Process a leased conversation job and send a minimal result to listen. @@ -146,8 +152,12 @@ async def send_result(result: Dict[str, Any]) -> None: data.extend(struct.pack("I", 201)) data.extend(bytes(json.dumps(result), "utf-8")) try: - await websocket.send_bytes(bytes(data)) - except (RuntimeError, WebSocketDisconnect): + if send_lock is None: + await websocket.send_bytes(bytes(data)) + else: + async with send_lock: + await websocket.send_bytes(bytes(data)) + except Exception: logger.info( 'pusher finalization result undeliverable after source close uid=%s conversation=%s', uid, @@ -317,7 +327,7 @@ async def _websocket_util_trigger( logger.info(f'_websocket_util_trigger {uid}') try: - await websocket.accept() + await websocket.accept(headers=[(b'x-omi-delivery-ack', b'1')]) except RuntimeError as e: logger.error(e) await websocket.close(code=1011, reason="Dirty state") @@ -338,8 +348,10 @@ async def _websocket_util_trigger( journey_attempt = JourneyAttempt('pusher_session') websocket_active = True shutdown_event = asyncio.Event() + websocket_send_lock = asyncio.Lock() websocket_close_code = 1000 application_failed = False + delivery_drain_requested = False try: # audio bytes @@ -363,12 +375,102 @@ async def _websocket_util_trigger( speaker_sample_queue: deque[_SpeakerSampleRequest] = deque(maxlen=SPEAKER_SAMPLE_QUEUE_WARN_SIZE) transcript_queue: deque[_TranscriptQueueItem] = deque(maxlen=TRANSCRIPT_QUEUE_WARN_SIZE) audio_bytes_queue: deque[_AudioBytesQueueItem] = deque(maxlen=AUDIO_BYTES_QUEUE_WARN_SIZE) + transcript_delivery_ids: set[str] = set() + speaker_sample_delivery_ids: set[str] = set() # private_cloud_queue caps at PRIVATE_CLOUD_QUEUE_MAX_SIZE to prevent OOM kills. # An OOM kill loses ALL queued data for ALL users on the pod — dropping the oldest # chunk for one user is strictly better than killing the pod. private_cloud_queue: deque[_PrivateCloudChunk] = deque(maxlen=PRIVATE_CLOUD_QUEUE_MAX_SIZE) audio_bytes_event = asyncio.Event() # Signals when items are added for instant wake + speaker_sample_event = asyncio.Event() + private_cloud_drained_event = asyncio.Event() + private_cloud_drained_event.set() + + def normalize_delivery_id(delivery_id: Any) -> Optional[str]: + if not isinstance(delivery_id, str) or not delivery_id or len(delivery_id) > 128: + return None + return delivery_id + + async def send_delivery_ack(kind: str, delivery_id: str) -> None: + data = bytearray() + data.extend(struct.pack('I', 202)) + data.extend(bytes(json.dumps({'kind': kind, 'delivery_id': delivery_id}), 'utf-8')) + try: + async with websocket_send_lock: + await websocket.send_bytes(bytes(data)) + except Exception: + logger.info('pusher delivery acknowledgement undeliverable uid=%s kind=%s', uid, kind) + + async def begin_delivery(kind: str, delivery_id: str) -> tuple[str, Optional[str]]: + state, lease_token = await run_blocking( + db_executor, + redis_db.begin_pusher_delivery, + uid, + f'{kind}:{delivery_id}', + ) + if state == 'unavailable': + record_fallback( + component='pusher', + from_mode='redis_delivery_lease', + to_mode='unfenced_delivery', + reason='other', + outcome='degraded', + log=logger, + ) + return state, lease_token + + async def finish_delivery( + kind: str, + delivery_id: str, + state: str, + lease_token: Optional[str], + ) -> None: + if state == 'claimed' and lease_token: + completed = await run_blocking( + db_executor, + redis_db.complete_pusher_delivery, + uid, + f'{kind}:{delivery_id}', + lease_token, + ) + if not completed: + record_fallback( + component='pusher', + from_mode='redis_delivery_completion', + to_mode='effect_acknowledgement', + reason='other', + outcome='degraded', + log=logger, + ) + await send_delivery_ack(kind, delivery_id) + + async def abandon_delivery( + kind: str, + delivery_id: Optional[str], + state: str, + lease_token: Optional[str], + ) -> None: + if not delivery_id or state != 'claimed' or not lease_token: + return + released = await asyncio.shield( + run_blocking( + db_executor, + redis_db.abandon_pusher_delivery, + uid, + f'{kind}:{delivery_id}', + lease_token, + ) + ) + if not released: + record_fallback( + component='pusher', + from_mode='redis_delivery_lease', + to_mode='lease_expiry_retry', + reason='other', + outcome='degraded', + log=logger, + ) async def process_private_cloud_queue() -> None: """Background task that batches private cloud sync uploads by conversation_id. @@ -379,6 +481,7 @@ async def process_private_cloud_queue() -> None: - The websocket disconnects (shutdown flush). """ nonlocal websocket_active + nonlocal delivery_drain_requested # Pending batches keyed by conversation_id pending: Dict[str, Dict[str, Any]] = {} @@ -460,47 +563,61 @@ async def _flush_batch(conv_id: str): ) del chunk_data - while websocket_active or len(private_cloud_queue) > 0 or len(pending) > 0: - await wait_for_event(shutdown_event, PRIVATE_CLOUD_SYNC_PROCESS_INTERVAL) - - # Drain queue into pending batches - if private_cloud_queue: - chunks_to_process = private_cloud_queue.copy() - private_cloud_queue.clear() - for chunk_info in chunks_to_process: - _add_to_batch(chunk_info) - - if not pending: - continue - - now = time.monotonic() - batch_size_threshold = sample_rate * 2 * PRIVATE_CLOUD_CHUNK_DURATION - - # Determine which conversations to flush - conv_ids_to_flush: List[str] = [] - for conv_id, batch in pending.items(): - batch_age = now - batch['queued_at'] - is_shutdown = not websocket_active - is_size_ready = len(batch['data']) >= batch_size_threshold - is_age_ready = batch_age >= PRIVATE_CLOUD_BATCH_MAX_AGE - if is_shutdown or is_size_ready or is_age_ready: - conv_ids_to_flush.append(conv_id) + try: + while websocket_active or len(private_cloud_queue) > 0 or len(pending) > 0: + await wait_for_event(shutdown_event, PRIVATE_CLOUD_SYNC_PROCESS_INTERVAL) + + # Drain queue into pending batches + if private_cloud_queue: + chunks_to_process = private_cloud_queue.copy() + private_cloud_queue.clear() + for chunk_info in chunks_to_process: + _add_to_batch(chunk_info) + + if not pending: + private_cloud_drained_event.set() + continue - for conv_id in conv_ids_to_flush: - await _flush_batch(conv_id) + now = time.monotonic() + batch_size_threshold = sample_rate * 2 * PRIVATE_CLOUD_CHUNK_DURATION + + # Determine which conversations to flush + conv_ids_to_flush: List[str] = [] + for conv_id, batch in pending.items(): + batch_age = now - batch['queued_at'] + is_shutdown = not websocket_active + is_size_ready = len(batch['data']) >= batch_size_threshold + is_age_ready = batch_age >= PRIVATE_CLOUD_BATCH_MAX_AGE + if delivery_drain_requested or is_shutdown or is_size_ready or is_age_ready: + conv_ids_to_flush.append(conv_id) + + for conv_id in conv_ids_to_flush: + await _flush_batch(conv_id) + if not private_cloud_queue and not pending: + private_cloud_drained_event.set() + finally: + if not private_cloud_queue and not pending: + private_cloud_drained_event.set() async def process_speaker_sample_queue() -> None: """Background task that processes speaker sample extraction requests.""" nonlocal websocket_active + nonlocal delivery_drain_requested while websocket_active or len(speaker_sample_queue) > 0: - await wait_for_event(shutdown_event, SPEAKER_SAMPLE_PROCESS_INTERVAL) + try: + await asyncio.wait_for(speaker_sample_event.wait(), timeout=SPEAKER_SAMPLE_PROCESS_INTERVAL) + except asyncio.TimeoutError: + pass + speaker_sample_event.clear() if not speaker_sample_queue: continue current_time = time.time() is_shutdown = not websocket_active + if is_shutdown or delivery_drain_requested: + await private_cloud_drained_event.wait() # Separate ready and pending requests. # On shutdown, skip the age check — process everything so pending @@ -509,7 +626,11 @@ async def process_speaker_sample_queue() -> None: pending_requests: List[_SpeakerSampleRequest] = [] for request in list(speaker_sample_queue): - if is_shutdown or current_time - request['queued_at'] >= SPEAKER_SAMPLE_MIN_AGE: + if ( + is_shutdown + or delivery_drain_requested + or current_time - request['queued_at'] >= SPEAKER_SAMPLE_MIN_AGE + ): ready_requests.append(request) else: pending_requests.append(request) @@ -523,17 +644,39 @@ async def process_speaker_sample_queue() -> None: person_id = request['person_id'] conv_id = request['conversation_id'] segment_ids = request['segment_ids'] + delivery_id = request['delivery_id'] + state = 'legacy' + lease_token: Optional[str] = None try: - await extract_speaker_samples( + if delivery_id: + state, lease_token = await begin_delivery('speaker_sample', delivery_id) + if state == 'done': + await send_delivery_ack('speaker_sample', cast(str, delivery_id)) + continue + if state == 'busy': + continue + result = await extract_speaker_samples( uid=uid, person_id=person_id, conversation_id=conv_id, segment_ids=segment_ids, sample_rate=sample_rate, + delivery_id=delivery_id, ) + if result.retryable: + raise RuntimeError(f'speaker sample extraction retryable: {result.reason}') + if delivery_id: + await finish_delivery('speaker_sample', delivery_id, state, lease_token) + except asyncio.CancelledError: + await abandon_delivery('speaker_sample', delivery_id, state, lease_token) + raise except Exception as e: + await abandon_delivery('speaker_sample', delivery_id, state, lease_token) logger.error(f"Error extracting speaker samples: {e} {uid} {conv_id}") + finally: + if delivery_id: + speaker_sample_delivery_ids.discard(delivery_id) async def process_transcript_queue() -> None: """Batched consumer for transcript events (realtime integrations + webhooks).""" @@ -552,11 +695,35 @@ async def process_transcript_queue() -> None: for item in batch: segments = item['segments'] memory_id = item['memory_id'] + delivery_id = item['delivery_id'] + state = 'legacy' + lease_token = None try: - await trigger_realtime_integrations(uid, segments, memory_id) - await realtime_transcript_webhook(uid, segments) + if delivery_id: + state, lease_token = await begin_delivery('transcript', delivery_id) + if state == 'done': + await send_delivery_ack('transcript', cast(str, delivery_id)) + continue + if state == 'busy': + continue + await trigger_realtime_integrations( + uid, + segments, + memory_id, + idempotency_key=delivery_id, + ) + await realtime_transcript_webhook(uid, segments, idempotency_key=delivery_id) + if delivery_id: + await finish_delivery('transcript', delivery_id, state, lease_token) + except asyncio.CancelledError: + await abandon_delivery('transcript', delivery_id, state, lease_token) + raise except Exception as e: + await abandon_delivery('transcript', delivery_id, state, lease_token) logger.error(f"Error processing transcript batch: {e} {uid}") + finally: + if delivery_id: + transcript_delivery_ids.discard(delivery_id) async def process_audio_bytes_queue() -> None: """Event-driven consumer for audio bytes triggers (app integrations + webhooks).""" @@ -593,6 +760,7 @@ async def receive_tasks() -> None: nonlocal websocket_active nonlocal websocket_close_code nonlocal application_failed + nonlocal delivery_drain_requested nonlocal speaker_sample_queue nonlocal transcript_queue nonlocal audio_bytes_queue @@ -643,6 +811,7 @@ async def receive_tasks() -> None: 'retries': 0, } ) + private_cloud_drained_event.clear() logger.info( f"Flushed private cloud buffer on conversation switch: {len(private_cloud_sync_buffer)} bytes {uid}" ) @@ -657,16 +826,69 @@ async def receive_tasks() -> None: res = json.loads(bytes(data[4:]).decode("utf-8")) segments = res.get('segments') memory_id = res.get('memory_id') + delivery_id = normalize_delivery_id(res.get('delivery_id')) + if not isinstance(segments, list): + logger.warning('Ignoring malformed transcript delivery uid=%s', uid) + continue + if delivery_id and delivery_id in transcript_delivery_ids: + continue # A transcript's memory_id must NOT overwrite the session's authoritative # current_conversation_id (which is set only by header 103). Doing so let a stale # lifecycle event carrying an older conversation's memory_id rebind a newer recording # session, mis-associating subsequent private-cloud audio (see issue #6952). if len(transcript_queue) >= TRANSCRIPT_QUEUE_WARN_SIZE: logger.warning(f"Warning: transcript_queue size {len(transcript_queue)} {uid}") + if delivery_id: + record_fallback( + component='pusher', + from_mode='transcript_delivery_queue', + to_mode='sender_retry', + reason='capacity_full', + outcome='degraded', + log=logger, + ) + # The deque has maxlen, so appending here would silently + # evict an accepted stable delivery. Reject the newest + # frame instead; negotiated senders retry stable frames. + continue # Route this transcript by its own memory_id when present, falling back to the # session's conversation id. This does not mutate session-scoped state. conversation_or_memory_id = memory_id or current_conversation_id - transcript_queue.append({'segments': segments, 'memory_id': conversation_or_memory_id}) + transcript_queue.append( + { + 'segments': segments, + 'memory_id': conversation_or_memory_id, + 'delivery_id': delivery_id, + } + ) + if delivery_id: + transcript_delivery_ids.add(delivery_id) + continue + + # Delivery drain request. The sender keeps this socket open while + # stable transcript and speaker work completes, so flush live + # audio metadata before speaker extraction without waiting for a + # disconnect that would make acknowledgements undeliverable. + if header_type == PUSHER_DELIVERY_DRAIN_OPCODE: + if private_cloud_sync_enabled and current_conversation_id and private_cloud_sync_buffer: + if len(private_cloud_queue) >= PRIVATE_CLOUD_QUEUE_MAX_SIZE: + logger.warning( + f"private_cloud_queue full ({len(private_cloud_queue)}/{PRIVATE_CLOUD_QUEUE_MAX_SIZE}), " + f"dropping oldest chunk to prevent OOM {uid}" + ) + private_cloud_queue.append( + { + 'data': bytes(private_cloud_sync_buffer), + 'conversation_id': current_conversation_id, + 'timestamp': private_cloud_chunk_start_time or time.time(), + 'retries': 0, + } + ) + private_cloud_sync_buffer = bytearray() + private_cloud_chunk_start_time = None + private_cloud_drained_event.clear() + delivery_drain_requested = True + speaker_sample_event.set() continue # Process conversation request @@ -692,6 +914,7 @@ async def receive_tasks() -> None: byok_keys, finalization_job_id if isinstance(finalization_job_id, str) else None, dispatch_generation if isinstance(dispatch_generation, int) else None, + websocket_send_lock, ), name=f'pusher_finalization:{uid}:{conversation_id}', ) @@ -703,9 +926,25 @@ async def receive_tasks() -> None: person_id = res.get('person_id') conv_id = res.get('conversation_id') segment_ids = res.get('segment_ids', []) + delivery_id = normalize_delivery_id(res.get('delivery_id')) if person_id and conv_id and segment_ids: + if delivery_id and delivery_id in speaker_sample_delivery_ids: + continue if len(speaker_sample_queue) >= SPEAKER_SAMPLE_QUEUE_WARN_SIZE: logger.warning(f"Warning: speaker_sample_queue size {len(speaker_sample_queue)} {uid}") + if delivery_id: + record_fallback( + component='pusher', + from_mode='speaker_sample_queue', + to_mode='sender_retry', + reason='capacity_full', + outcome='degraded', + log=logger, + ) + # Preserve accepted FIFO work. A maxlen append would + # evict the oldest request without completing or + # acknowledging its stable delivery. + continue logger.info( f"Queued speaker sample request: person={person_id}, {len(segment_ids)} segments {uid}" ) @@ -715,12 +954,20 @@ async def receive_tasks() -> None: 'conversation_id': conv_id, 'segment_ids': segment_ids, 'queued_at': time.time(), + 'delivery_id': delivery_id, } ) + speaker_sample_event.set() + if delivery_id: + speaker_sample_delivery_ids.add(delivery_id) continue # Audio bytes if header_type == 101: + audio_conversation_id = current_conversation_id + if len(data) < 12: + logger.warning('Ignoring malformed audio delivery uid=%s', uid) + continue # Parse: header(4) | timestamp(8 bytes double) | audio_data buffer_start_timestamp = struct.unpack("d", data[4:12])[0] audio_data = data[12:] @@ -733,7 +980,7 @@ async def receive_tasks() -> None: audiobuffer.extend(audio_data) # Private cloud sync - queue chunks for background processing - if private_cloud_sync_enabled and current_conversation_id: + if private_cloud_sync_enabled and audio_conversation_id: if private_cloud_chunk_start_time is None: # Use timestamp from first buffer of this 5-second chunk private_cloud_chunk_start_time = buffer_start_timestamp @@ -749,11 +996,12 @@ async def receive_tasks() -> None: private_cloud_queue.append( { 'data': bytes(private_cloud_sync_buffer), - 'conversation_id': current_conversation_id, + 'conversation_id': audio_conversation_id, 'timestamp': cast(float, private_cloud_chunk_start_time), 'retries': 0, } ) + private_cloud_drained_event.clear() private_cloud_sync_buffer = bytearray() private_cloud_chunk_start_time = None @@ -813,8 +1061,10 @@ async def receive_tasks() -> None: 'retries': 0, } ) + private_cloud_drained_event.clear() logger.info(f"Flushed final private cloud buffer: {len(private_cloud_sync_buffer)} bytes {uid}") websocket_active = False + speaker_sample_event.set() bg_main_tasks: List[asyncio.Task[Any]] = [] try: @@ -847,6 +1097,7 @@ async def receive_tasks() -> None: if not receive_task.done(): websocket_active = False + speaker_sample_event.set() receive_task.cancel() try: await receive_task @@ -866,6 +1117,7 @@ async def receive_tasks() -> None: finally: shutdown_event.set() websocket_active = False + speaker_sample_event.set() all_to_cancel = [t for t in bg_main_tasks if not t.done()] await drain_tasks(all_to_cancel, timeout=5.0, label="pusher_cleanup", cancel=True) diff --git a/backend/routers/task_integrations.py b/backend/routers/task_integrations.py index e2d14e88f60..2ff21009abc 100644 --- a/backend/routers/task_integrations.py +++ b/backend/routers/task_integrations.py @@ -343,6 +343,9 @@ class CreateTaskResponse(BaseModel): success: bool external_task_id: Optional[str] = None error: Optional[str] = None + error_code: Optional[str] = None + retryable: Optional[bool] = None + ambiguous: Optional[bool] = None @router.post("/v1/task-integrations/{app_key}/tasks", response_model=CreateTaskResponse, tags=['task-integrations']) @@ -391,6 +394,9 @@ async def create_task_via_integration( success=result.get("success", False), external_task_id=result.get("external_task_id"), error=result.get("error"), + error_code=result.get("error_code"), + retryable=result.get("retryable"), + ambiguous=result.get("ambiguous"), ) diff --git a/backend/routers/users.py b/backend/routers/users.py index 97ca5be98af..c477210b99e 100644 --- a/backend/routers/users.py +++ b/backend/routers/users.py @@ -26,7 +26,6 @@ from services.users.data_export import iter_user_data_export from services.users.account_deletion import background_wipe_user_data, start_account_deletion from database.app_review_config import should_hide_subscription_ui -from database.webhook_health import record_dev_webhook_success from database.conversations import get_in_progress_conversation, get_conversation from database.redis_db import ( cache_user_geolocation, @@ -120,7 +119,7 @@ delete_user_person_speech_samples, delete_user_person_speech_sample, ) -from utils.webhooks import webhook_first_time_setup +from utils.webhooks import reset_user_webhook_delivery_health, webhook_first_time_setup from utils.byok import has_byok_keys, invalidate_byok_state_cache, peppered_fingerprint import logging @@ -472,10 +471,11 @@ class SetUserWebhookUrlRequest(BaseModel): def set_user_webhook_endpoint( wtype: WebhookType, data: SetUserWebhookUrlRequest, uid: str = Depends(auth.get_current_user_uid) ): - url = data.url - if url == '' or url == ',': + set_user_webhook_db(uid, wtype, data.url) + if data.url == '' or data.url == ',': disable_user_webhook_db(uid, wtype) - set_user_webhook_db(uid, wtype, url) + else: + enable_user_webhook_endpoint(wtype, uid) return {'status': 'ok'} @@ -493,7 +493,7 @@ def disable_user_webhook_endpoint(wtype: WebhookType, uid: str = Depends(auth.ge @router.post('/v1/users/developer/webhook/{wtype}/enable', tags=['v1'], response_model=UserStatusResponse) def enable_user_webhook_endpoint(wtype: WebhookType, uid: str = Depends(auth.get_current_user_uid)): enable_user_webhook_db(uid, wtype) - record_dev_webhook_success(uid, wtype.value) + reset_user_webhook_delivery_health(uid, wtype, get_user_webhook_db(uid, wtype)) return {'status': 'ok'} diff --git a/backend/testing/e2e/test_task_integrations.py b/backend/testing/e2e/test_task_integrations.py index 9476f9a8283..c14523c25a7 100644 --- a/backend/testing/e2e/test_task_integrations.py +++ b/backend/testing/e2e/test_task_integrations.py @@ -72,7 +72,14 @@ def handler(request): _close_async_client(fake_client) assert created.status_code == 200, created.text - assert created.json() == {"success": True, "external_task_id": "todo-123", "error": None} + assert created.json() == { + "success": True, + "external_task_id": "todo-123", + "error": None, + "error_code": None, + "retryable": None, + "ambiguous": None, + } assert len(requests) == 1 request = requests[0] assert str(request.url) == "https://api.todoist.com/rest/v2/tasks" @@ -145,6 +152,9 @@ def handler(request): "success": False, "external_task_id": None, "error": "Todoist API error: 500", + "error_code": "api_error", + "retryable": False, + "ambiguous": True, } assert len(requests) == 1 assert _get_todoist_integration(client, auth_headers)["connected"] is True @@ -196,5 +206,8 @@ def handler(request): assert response.json() == { "success": False, "external_task_id": None, - "error": "deterministic Todoist timeout", + "error": "ConnectTimeout", + "error_code": "transport_error", + "retryable": True, + "ambiguous": False, } diff --git a/backend/testing/e2e/test_webhooks.py b/backend/testing/e2e/test_webhooks.py index b6a4a7a05d3..7fd0b75a1e1 100644 --- a/backend/testing/e2e/test_webhooks.py +++ b/backend/testing/e2e/test_webhooks.py @@ -74,6 +74,9 @@ async def handler(request): def test_realtime_webhook_does_not_call_provider_when_disabled(client, auth_headers, monkeypatch, fake_redis): _configure_realtime_webhook(client, auth_headers) + disabled = client.post("/v1/users/developer/webhook/realtime_transcript/disable", headers=auth_headers) + assert disabled.status_code == 200, disabled.text + health_before = _health(fake_redis) requests = [] async def handler(request): @@ -83,7 +86,7 @@ async def handler(request): _run_realtime_delivery(monkeypatch, handler) assert requests == [] - assert _health(fake_redis) == {} + assert _health(fake_redis) == health_before @pytest.mark.parametrize( diff --git a/backend/testing/workflow_contracts.json b/backend/testing/workflow_contracts.json index 57edf09aa42..d337a373a4f 100644 --- a/backend/testing/workflow_contracts.json +++ b/backend/testing/workflow_contracts.json @@ -358,12 +358,23 @@ { "id": "webhook_delivery_health", "risk": "high", - "sources": ["backend/utils/webhooks.py", "backend/database/webhook_health.py"], - "tests": ["tests/unit/test_async_webhooks.py", "tests/unit/test_webhook_auto_disable.py"], + "sources": [ + "backend/utils/webhooks.py", + "backend/utils/http_client.py", + "backend/database/webhook_health.py", + "backend/routers/users.py" + ], + "tests": [ + "tests/unit/test_async_webhooks.py", + "tests/unit/test_async_http_infrastructure.py", + "tests/unit/test_webhook_auto_disable.py", + "tests/unit/test_users_webhook_url_validation.py" + ], "checks": [], "invariants": [ "delivery failures are counted durably", - "auto-disable state is endpoint-aware and race-safe" + "auto-disable state is endpoint-aware and race-safe", + "manual enable resets persistent and process-local failure gates without recording a synthetic delivery success" ] }, { @@ -456,6 +467,37 @@ "discard and generic status-write escape hatches remain characterized until the lifecycle service replaces them" ] }, + { + "id": "legacy_memory_projection_convergence", + "risk": "high", + "sources": [ + "backend/database/memories.py", + "backend/database/memory_ledger.py", + "backend/database/projection_repair.py", + "backend/database/vector_db.py", + "backend/routers/memories.py", + "backend/utils/memory/memory_service.py", + "backend/utils/retrieval/tool_services/memories.py", + "backend/utils/x_connector.py" + ], + "tests": [ + "tests/unit/test_memory_service_parity.py", + "tests/unit/test_memories_batch_delete.py", + "tests/unit/test_memories_delete_batch_chunk.py", + "tests/unit/test_memory_ledger.py", + "tests/unit/test_memories_batch.py", + "tests/unit/test_mcp_search_memories.py", + "tests/unit/test_lock_bypass_fixes.py", + "tests/unit/test_tools_rest_memory_runtime_adapter.py", + "tests/unit/test_x_memory_extraction_retry.py" + ], + "checks": ["no_large_tuple_results"], + "invariants": [ + "legacy mutation guards, reads, and writes use one selected Firestore client", + "deletions enumerate every authoritative memory id before removing projections", + "search overfetches and filters stale, locked, rejected, invalid, or invalidated projections" + ] + }, { "id": "listen_pusher_pipeline", "risk": "high", @@ -463,8 +505,12 @@ "backend/routers/transcribe.py", "backend/routers/listen/**", "backend/routers/pusher.py", + "backend/database/redis_db.py", "backend/utils/listen_pusher_session.py", + "backend/utils/speaker_identification.py", "backend/utils/pusher.py", + "backend/utils/app_integrations.py", + "backend/utils/webhooks.py", "backend/utils/stt/live_failure.py", "backend/utils/stt/streaming.py", "backend/config/stt_provider_policy.py", @@ -481,6 +527,12 @@ "tests": [ "tests/unit/test_listen_pipeline.py", "tests/unit/test_pusher_heartbeat.py", + "tests/unit/test_pusher_readiness_drain.py", + "tests/unit/test_redis_db_cache_serialization.py", + "tests/unit/test_speaker_identification_delivery.py", + "tests/unit/test_listen_finalization_cloud_tasks.py", + "tests/unit/test_async_app_integrations.py", + "tests/unit/test_async_webhooks.py", "tests/unit/test_pusher_conversation_retry.py", "tests/unit/test_live_stt_failure.py", "tests/unit/test_listen_runtime_regressions.py", @@ -498,6 +550,13 @@ "teardown flushes tail audio before pusher close", "pending conversation requests retry until ack or give-up limit", "pusher heartbeat prevents idle connection drops", + "draining pusher pods reject new sockets cleanly while preserving established-session shutdown", + "transcript and speaker deliveries retain stable identity until a local-invocation completion acknowledgement on negotiated sockets", + "route-stamped audio survives explicit send failure and cancellation while ambiguous local completion remains documented as best effort", + "pusher leases stable deliveries only when a worker owns the effect and never evicts an accepted unacknowledged frame", + "done markers suppress stable delivery replays across sockets and replicas while Redis failure remains observable and fail-open", + "graceful close performs a bounded acknowledgement drain after the listen runtime becomes inactive", + "downstream realtime webhooks receive stable idempotency keys while opcode 202 acknowledges local worker invocation rather than downstream HTTP delivery", "a client Parakeet preference cannot route a live session to an incompatible model", "listen-to-pusher boundary changes run the credential-free local Firestore, Redis, and Parakeet-stub stack gauntlet in PR CI" ] diff --git a/backend/tests/unit/test_async_app_integrations.py b/backend/tests/unit/test_async_app_integrations.py index b7df57d283d..1332c167520 100644 --- a/backend/tests/unit/test_async_app_integrations.py +++ b/backend/tests/unit/test_async_app_integrations.py @@ -11,6 +11,8 @@ import pytest +from testing.import_isolation import load_module_fresh + os.environ.setdefault( "ENCRYPTION_SECRET", "omi_ZwB2ZNqB2HHpMK6wStk7sTpavJiPTFg7gXUHnc4tFABPU6pZ2c2DKgehtfgi4RZv", @@ -269,9 +271,10 @@ async def _run_blocking(_executor, func, *args, **kwargs): _executors_mod.run_blocking = _run_blocking -import importlib - -app_integrations = importlib.import_module("utils.app_integrations") +app_integrations = load_module_fresh( + "utils.app_integrations", + os.path.join(_BACKEND_DIR, "utils", "app_integrations.py"), +) _restore_stub_modules() @@ -564,6 +567,72 @@ async def test_12_apps_sent_in_two_chunks(self): class TestAsyncTriggerRealtimeIntegrations: """Test async realtime integration fan-out.""" + @pytest.mark.asyncio + async def test_retryable_status_reuses_one_receiver_visible_delivery_key(self): + unavailable = MagicMock(status_code=503) + accepted = MagicMock(status_code=204) + client = AsyncMock() + client.post = AsyncMock(side_effect=[unavailable, accepted]) + + with patch.object(app_integrations, "get_webhook_client", return_value=client): + result = await app_integrations._post_realtime_app_webhook( + "app-1", + "https://app.test/hook", + idempotency_key="delivery-1", + retry_delays=(0,), + json={"segments": [{"text": "hi"}]}, + ) + + assert result is accepted + assert client.post.await_count == 2 + assert [call.kwargs["headers"]["X-Omi-Idempotency-Key"] for call in client.post.await_args_list] == [ + "delivery-1", + "delivery-1", + ] + + @pytest.mark.asyncio + async def test_transport_retry_generates_one_stable_key_for_legacy_caller(self): + accepted = MagicMock(status_code=200) + client = AsyncMock() + client.post = AsyncMock( + side_effect=[ + app_integrations.httpx.ConnectError("connection unavailable"), + accepted, + ] + ) + + with patch.object(app_integrations, "get_webhook_client", return_value=client): + result = await app_integrations._post_realtime_app_webhook( + "app-1", + "https://app.test/hook", + retry_delays=(0,), + json={"segments": [{"text": "hi"}]}, + ) + + assert result is accepted + keys = [call.kwargs["headers"]["X-Omi-Idempotency-Key"] for call in client.post.await_args_list] + assert len(keys) == 2 + assert keys[0] == keys[1] + assert keys[0] + + @pytest.mark.asyncio + async def test_permanent_client_error_is_not_retried(self): + rejected = MagicMock(status_code=400) + client = AsyncMock() + client.post = AsyncMock(return_value=rejected) + + with patch.object(app_integrations, "get_webhook_client", return_value=client): + result = await app_integrations._post_realtime_app_webhook( + "app-1", + "https://app.test/hook", + idempotency_key="delivery-1", + retry_delays=(0, 0), + json={"segments": [{"text": "hi"}]}, + ) + + assert result is rejected + client.post.assert_awaited_once() + @pytest.mark.asyncio async def test_no_apps_returns_empty(self): """No apps and no mentor → empty result.""" @@ -594,6 +663,26 @@ async def test_multiple_apps_called_concurrently(self): assert mock_client.post.call_count == 2 + @pytest.mark.asyncio + async def test_stable_delivery_id_reaches_realtime_app_webhook(self): + app = _make_app("a1", "https://app1.test/hook", triggers_realtime=True) + response = MagicMock(status_code=200, text="") + response.json.return_value = {} + client = AsyncMock() + client.post = AsyncMock(return_value=response) + + with patch.object(app_integrations, "get_available_apps", return_value=[app]), patch.object( + app_integrations, "process_mentor_notification", return_value=None + ), patch.object(app_integrations, "get_webhook_client", return_value=client): + await app_integrations.trigger_realtime_integrations( + "uid-1", + [{"text": "hi"}], + "conv-1", + idempotency_key="delivery-1", + ) + + assert client.post.await_args.kwargs["headers"] == {"X-Omi-Idempotency-Key": "delivery-1"} + @pytest.mark.asyncio async def test_app_response_message_triggers_notification(self): """App returning a message > 5 chars triggers notification.""" diff --git a/backend/tests/unit/test_async_http_infrastructure.py b/backend/tests/unit/test_async_http_infrastructure.py index c9eee2f348d..764e7a4449c 100644 --- a/backend/tests/unit/test_async_http_infrastructure.py +++ b/backend/tests/unit/test_async_http_infrastructure.py @@ -56,6 +56,7 @@ def _drop_stale_module(name, required_attrs): _drop_stale_module("utils.http_client", ["WebhookCircuitBreaker", "get_webhook_circuit_breaker"]) _drop_stale_module("utils.executors", ["critical_executor", "storage_executor", "shutdown_executors"]) +import utils.http_client as http_client_module from utils.http_client import ( WebhookCircuitBreaker, get_webhook_circuit_breaker, @@ -71,6 +72,7 @@ def _drop_stale_module(name, required_attrs): _SEMAPHORE_CACHE_MAX, _CIRCUIT_BREAKER_FAILURE_THRESHOLD, _CIRCUIT_BREAKER_RECOVERY_TIMEOUT, + reset_webhook_circuit_breaker, ) from utils.executors import critical_executor, storage_executor @@ -219,6 +221,18 @@ def test_invalid_url_fallback(self): assert cb is not None assert cb.state == 'closed' + def test_same_path_url_replacement_is_allowed_immediately(self): + old_cb = get_webhook_circuit_breaker("https://example.com/hook?version=old") + for _ in range(_CIRCUIT_BREAKER_FAILURE_THRESHOLD): + old_cb.record_failure() + assert old_cb.allow_request() is False + + reset_webhook_circuit_breaker("https://example.com/hook?version=new") + + replacement_cb = get_webhook_circuit_breaker("https://example.com/hook?version=new") + assert replacement_cb is not old_cb + assert replacement_cb.allow_request() is True + # ============================================================================ # Latest-wins dropping @@ -571,7 +585,7 @@ def test_queue_max_size_is_20(self): pytest.fail("PRIVATE_CLOUD_QUEUE_MAX_SIZE constant not found") def test_overflow_warning_at_all_enqueue_points(self): - """All 3 enqueue points must log overflow warning before deque drops oldest.""" + """All enqueue points must log overflow warning before deque drops oldest.""" import os backend_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) @@ -580,7 +594,7 @@ def test_overflow_warning_at_all_enqueue_points(self): # Count occurrences of the overflow warning pattern warning_count = src.count('private_cloud_queue full') - assert warning_count == 3, f"Expected 3 overflow warnings, found {warning_count}" + assert warning_count == 4, f"Expected 4 overflow warnings, found {warning_count}" def test_deque_maxlen_drops_oldest(self): """Verify deque(maxlen=N) drops oldest item when full.""" @@ -638,15 +652,12 @@ def test_stale_breaker_evicted(self): assert 'https://stale.test/hook' not in _webhook_circuit_breakers _webhook_circuit_breakers.clear() - def test_allow_request_updates_access_time(self): + def test_allow_request_updates_access_time(self, monkeypatch): """allow_request() must update _last_access_time.""" - import time - from utils.http_client import _webhook_circuit_breakers, get_webhook_circuit_breaker - _webhook_circuit_breakers.clear() cb = get_webhook_circuit_breaker('https://test.test/hook') - old_access = cb._last_access_time - time.sleep(0.01) + cb._last_access_time = 100.0 + monkeypatch.setattr(http_client_module.time, 'monotonic', lambda: 101.0) cb.allow_request() - assert cb._last_access_time > old_access + assert cb._last_access_time == 101.0 _webhook_circuit_breakers.clear() diff --git a/backend/tests/unit/test_async_webhooks.py b/backend/tests/unit/test_async_webhooks.py index 6b2176ede8e..262b280f53c 100644 --- a/backend/tests/unit/test_async_webhooks.py +++ b/backend/tests/unit/test_async_webhooks.py @@ -9,10 +9,15 @@ import re from unittest.mock import MagicMock, AsyncMock, patch +import httpx import pytest import utils.webhooks as webhooks_module -from utils.webhooks import realtime_transcript_webhook, send_audio_bytes_developer_webhook, day_summary_webhook +from utils.webhooks import ( + day_summary_webhook, + realtime_transcript_webhook, + send_audio_bytes_developer_webhook, +) @pytest.fixture(autouse=True) @@ -31,6 +36,82 @@ def _stub_webhook_db_helpers(monkeypatch): monkeypatch.setattr(webhooks_module, "record_dev_webhook_failure", MagicMock(return_value=False)) +class TestPostDevWebhookRetryPolicy: + @pytest.mark.asyncio + @pytest.mark.parametrize('status_code', [400, 401, 403, 404, 409, 410, 422]) + async def test_permanent_4xx_is_not_retried(self, status_code): + response = MagicMock(status_code=status_code) + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=response) + mock_sleep = AsyncMock() + + with patch.object(webhooks_module, 'get_webhook_client', return_value=mock_client), patch.object( + webhooks_module.asyncio, 'sleep', new=mock_sleep + ): + actual = await webhooks_module._post_dev_webhook( + 'test_webhook', + 'https://example.com/webhook', + retry_delays=(1.0,), + json={'event': 'conversation.completed'}, + ) + + assert actual is response + mock_client.post.assert_awaited_once() + mock_sleep.assert_not_awaited() + + @pytest.mark.asyncio + @pytest.mark.parametrize('status_code', [408, 425, 429, 500, 503]) + async def test_retryable_http_status_retries_with_stable_idempotency_key(self, status_code): + retryable_response = MagicMock(status_code=status_code) + success_response = MagicMock(status_code=204) + mock_client = AsyncMock() + mock_client.post = AsyncMock(side_effect=[retryable_response, success_response]) + mock_sleep = AsyncMock() + + with patch.object(webhooks_module, 'get_webhook_client', return_value=mock_client), patch.object( + webhooks_module.asyncio, 'sleep', new=mock_sleep + ): + actual = await webhooks_module._post_dev_webhook( + 'test_webhook', + 'https://example.com/webhook', + retry_delays=(1.0,), + json={'event': 'conversation.completed'}, + ) + + assert actual is success_response + assert mock_client.post.await_count == 2 + mock_sleep.assert_awaited_once_with(1.0) + idempotency_keys = [call.kwargs['headers']['Idempotency-Key'] for call in mock_client.post.await_args_list] + assert idempotency_keys[0] + assert idempotency_keys[0] == idempotency_keys[1] + + @pytest.mark.asyncio + async def test_network_failure_retries_with_stable_idempotency_key(self): + request = httpx.Request('POST', 'https://example.com/webhook') + network_error = httpx.ConnectError('connection failed', request=request) + success_response = MagicMock(status_code=200) + mock_client = AsyncMock() + mock_client.post = AsyncMock(side_effect=[network_error, success_response]) + mock_sleep = AsyncMock() + + with patch.object(webhooks_module, 'get_webhook_client', return_value=mock_client), patch.object( + webhooks_module.asyncio, 'sleep', new=mock_sleep + ): + actual = await webhooks_module._post_dev_webhook( + 'test_webhook', + 'https://example.com/webhook', + retry_delays=(1.0,), + json={'event': 'conversation.completed'}, + ) + + assert actual is success_response + assert mock_client.post.await_count == 2 + mock_sleep.assert_awaited_once_with(1.0) + idempotency_keys = [call.kwargs['headers']['Idempotency-Key'] for call in mock_client.post.await_args_list] + assert idempotency_keys[0] + assert idempotency_keys[0] == idempotency_keys[1] + + class TestRealtimeTranscriptWebhook: """Test realtime_transcript_webhook uses httpx async.""" @@ -51,6 +132,21 @@ async def test_success_sends_via_httpx(self): call_args = mock_client.post.call_args assert "segments" in call_args.kwargs.get("json", {}) + @pytest.mark.asyncio + async def test_stable_delivery_id_reaches_realtime_webhook(self): + mock_response = MagicMock(status_code=204) + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=mock_response) + + with patch.object(webhooks_module, "get_webhook_client", return_value=mock_client): + await realtime_transcript_webhook( + "uid-1", + [{"text": "hello"}], + idempotency_key="delivery-1", + ) + + assert mock_client.post.await_args.kwargs["headers"]["Idempotency-Key"] == "delivery-1" + @pytest.mark.asyncio async def test_notification_on_200_with_message(self): """Verify webhook notification sent when response has message > 5 chars.""" @@ -103,7 +199,7 @@ async def test_timeout_error_handled(self): mock_client.post = AsyncMock(side_effect=httpx.TimeoutException("connect timeout")) with patch("utils.webhooks.get_webhook_client", return_value=mock_client), patch( - "utils.webhooks._get_dev_webhook_retry_delays", return_value=() + "utils.webhooks._REALTIME_DEV_WEBHOOK_RETRY_DELAYS", () ): # Should not raise await realtime_transcript_webhook("uid-1", [{"text": "hello"}]) @@ -396,7 +492,7 @@ async def test_transcript_webhook_records_failure_on_exception(self): with patch("utils.webhooks.get_webhook_circuit_breaker", return_value=mock_cb), patch( "utils.webhooks.get_webhook_client", return_value=mock_client - ), patch("utils.webhooks._get_dev_webhook_retry_delays", return_value=()): + ), patch("utils.webhooks._REALTIME_DEV_WEBHOOK_RETRY_DELAYS", ()): await realtime_transcript_webhook("uid-1", [{"text": "hello"}]) mock_cb.record_failure.assert_called_once() diff --git a/backend/tests/unit/test_dev_api_lock_bypass.py b/backend/tests/unit/test_dev_api_lock_bypass.py index 71a8989cfdc..97ff94d62c6 100644 --- a/backend/tests/unit/test_dev_api_lock_bypass.py +++ b/backend/tests/unit/test_dev_api_lock_bypass.py @@ -396,15 +396,22 @@ def test_delete_memory_allows_unlocked(self): import database.memories as memories_db memories_db.get_memory = MagicMock(return_value=_make_memory(locked=False)) - memories_db.delete_memory = MagicMock() + memories_db.delete_memory = MagicMock(return_value=SimpleNamespace(committed_count=1)) - from routers.developer import delete_memory + from routers import developer as developer_module _allow_developer_memory_write_grant() - result = delete_memory(memory_id='mem-1', auth_context=_developer_memory_write_context()) + result = developer_module.delete_memory( + memory_id='mem-1', + auth_context=_developer_memory_write_context(), + ) assert result == {"success": True} - memories_db.delete_memory.assert_called_once_with('test-uid', 'mem-1') + memories_db.delete_memory.assert_called_once_with( + 'test-uid', + 'mem-1', + firestore_client=developer_module.db, + ) # ============================================================================= diff --git a/backend/tests/unit/test_developer_memory_adapter.py b/backend/tests/unit/test_developer_memory_adapter.py index 3bc8f41d9db..f42a1856b1a 100644 --- a/backend/tests/unit/test_developer_memory_adapter.py +++ b/backend/tests/unit/test_developer_memory_adapter.py @@ -59,6 +59,28 @@ def _function_source_for_route(path: str, method: str) -> str: raise AssertionError(f'route not found: {method.upper()} {path}') +def _class_method_ast(path: Path, class_name: str, method_name: str) -> ast.FunctionDef: + module = ast.parse(path.read_text(encoding='utf-8')) + for node in module.body: + if not isinstance(node, ast.ClassDef) or node.name != class_name: + continue + for child in node.body: + if isinstance(child, ast.FunctionDef) and child.name == method_name: + return child + raise AssertionError(f'method not found: {class_name}.{method_name}') + + +def _qualified_call_name(call: ast.Call) -> str: + parts = [] + node = call.func + while isinstance(node, ast.Attribute): + parts.append(node.attr) + node = node.value + if isinstance(node, ast.Name): + parts.append(node.id) + return '.'.join(reversed(parts)) + + def _memory_item(memory_id: str, *, tier=MemoryTier.short_term, now=None, captured_at=None, content=None, **overrides): return memory_item( memory_id, @@ -151,20 +173,29 @@ def test_developer_batch_create_route_checks_split_brain_guard_before_categoriza def test_developer_delete_route_checks_split_brain_guard_before_reads_and_legacy_delete(): memory_service_py = Path(__file__).resolve().parents[2] / 'utils' / 'memory' / 'memory_service.py' route_source = _function_source_for_route('/v1/dev/user/memories/{memory_id}', 'delete') - service_contents = memory_service_py.read_text(encoding='utf-8') + method = _class_method_ast(memory_service_py, 'MemoryService', 'delete_external_memory') pin_call = 'pin_memory_system(uid, db_client=db)' external_delete = '.delete_external_memory(' - guard_call = 'guard_legacy_memory_write(' - legacy_read = 'memory = memories_db.get_memory(uid, memory_id)' - legacy_delete = 'memories_db.delete_memory(uid, memory_id)' + calls = sorted( + (node for node in ast.walk(method) if isinstance(node, ast.Call)), + key=lambda node: (node.lineno, node.col_offset), + ) + calls_by_name = {_qualified_call_name(call): call for call in calls} + guard_call = calls_by_name['_require_legacy_write_guard'] + legacy_read = calls_by_name['memories_db.get_memory'] + legacy_delete = calls_by_name['memories_db.delete_memory'] assert pin_call in route_source assert external_delete in route_source - assert guard_call in service_contents - assert legacy_read in service_contents - assert legacy_delete in service_contents assert route_source.index(pin_call) < route_source.index(external_delete) - guard_index = service_contents.index(guard_call) - assert guard_index < service_contents.index(legacy_read) < service_contents.index(legacy_delete) + assert guard_call.lineno < legacy_read.lineno < legacy_delete.lineno + for call in (legacy_read, legacy_delete): + firestore_client = next((keyword.value for keyword in call.keywords if keyword.arg == 'firestore_client'), None) + assert ( + isinstance(firestore_client, ast.Attribute) + and firestore_client.attr == '_db_client' + and isinstance(firestore_client.value, ast.Name) + and firestore_client.value.id == 'self' + ), f'{_qualified_call_name(call)} must use the injected Firestore client' def test_developer_update_route_checks_split_brain_guard_before_reads_and_legacy_mutations(): diff --git a/backend/tests/unit/test_listen_finalization_cloud_tasks.py b/backend/tests/unit/test_listen_finalization_cloud_tasks.py index f6134435fb2..8ed924d25cf 100644 --- a/backend/tests/unit/test_listen_finalization_cloud_tasks.py +++ b/backend/tests/unit/test_listen_finalization_cloud_tasks.py @@ -6,6 +6,7 @@ import json from pathlib import Path import runpy +import struct from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch @@ -24,6 +25,7 @@ from utils import cloud_tasks from utils.conversations.finalizer import ConversationFinalizationDisposition, ConversationFinalizationError import utils.conversations.finalizer as persisted_finalizer +from utils.speaker_identification import SpeakerSampleExtractionResult def _prod_backend_sync_runtime_env(monkeypatch): @@ -441,14 +443,20 @@ class _PusherLifecycleWebSocket: def __init__(self, receive_bytes): self._receive_bytes = receive_bytes self.accepted = False + self.accept_headers = None + self.sent: list[bytes] = [] self.client_state = pusher_router.WebSocketState.DISCONNECTED - async def accept(self) -> None: + async def accept(self, *, headers=None) -> None: self.accepted = True + self.accept_headers = headers async def receive_bytes(self) -> bytes: return await self._receive_bytes() + async def send_bytes(self, payload: bytes) -> None: + self.sent.append(payload) + class _PusherJourneyAttempt: outcomes: list[str] = [] @@ -484,6 +492,41 @@ async def _inline_run_blocking(_executor, func, *args, **kwargs): return func(*args, **kwargs) +def _patch_pusher_worker_drain(monkeypatch) -> None: + async def supervisor(*, receive_task, **_kwargs): + await receive_task + return SimpleNamespace(reason='disconnect', task_name='ws:uid-1:receive') + + async def drain(tasks, *, cancel, **_kwargs): + tasks = list(tasks) + if cancel: + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + + monkeypatch.setattr(pusher_router, 'supervise_tasks', supervisor) + monkeypatch.setattr(pusher_router, 'drain_tasks', drain) + + +def _transcript_delivery_frame(delivery_id: str, *, segment_id: str = 'segment-1') -> bytes: + payload = { + 'segments': [{'id': segment_id, 'text': 'hello'}], + 'memory_id': 'conversation-1', + 'delivery_id': delivery_id, + } + return struct.pack('I', 102) + json.dumps(payload).encode('utf-8') + + +def _speaker_delivery_frame(delivery_id: str) -> bytes: + payload = { + 'person_id': 'person-1', + 'conversation_id': 'conversation-1', + 'segment_ids': ['segment-1'], + 'delivery_id': delivery_id, + } + return struct.pack('I', 105) + json.dumps(payload).encode('utf-8') + + @pytest.mark.anyio async def test_worker_retries_processing_failure_before_final_attempt(monkeypatch): monkeypatch.setattr(finalization_router, 'run_blocking', _inline_run_blocking) @@ -762,6 +805,295 @@ async def supervisor(*, receive_task, **_kwargs): assert _PusherJourneyAttempt.outcomes == ['failure'] +@pytest.mark.anyio +async def test_pusher_acknowledges_transcript_only_after_owned_effects_complete(monkeypatch): + _PusherJourneyAttempt.outcomes = [] + _patch_pusher_session_dependencies(monkeypatch) + _patch_pusher_worker_drain(monkeypatch) + receive_bytes = AsyncMock( + side_effect=[ + _transcript_delivery_frame('delivery-1'), + pusher_router.WebSocketDisconnect(code=1000), + ] + ) + websocket = _PusherLifecycleWebSocket(receive_bytes=receive_bytes) + events: list[str] = [] + begin = MagicMock(side_effect=lambda *_args: events.append('begin') or ('claimed', 'lease-1')) + complete = MagicMock(side_effect=lambda *_args: events.append('complete') or True) + + async def trigger(*_args, **_kwargs): + events.append('app') + + async def webhook(*_args, **_kwargs): + events.append('webhook') + + monkeypatch.setattr(pusher_router.redis_db, 'begin_pusher_delivery', begin) + monkeypatch.setattr(pusher_router.redis_db, 'complete_pusher_delivery', complete) + monkeypatch.setattr(pusher_router, 'trigger_realtime_integrations', trigger) + monkeypatch.setattr(pusher_router, 'realtime_transcript_webhook', webhook) + + await pusher_router._websocket_util_trigger(websocket, 'uid-1') + + assert websocket.accept_headers == [(b'x-omi-delivery-ack', b'1')] + assert events == ['begin', 'app', 'webhook', 'complete'] + assert len(websocket.sent) == 1 + assert int.from_bytes(websocket.sent[0][:4], 'little') == 202 + assert json.loads(websocket.sent[0][4:]) == { + 'kind': 'transcript', + 'delivery_id': 'delivery-1', + } + + +@pytest.mark.anyio +async def test_pusher_unhandled_transcript_worker_failure_leaves_delivery_unacknowledged(monkeypatch): + _PusherJourneyAttempt.outcomes = [] + _patch_pusher_session_dependencies(monkeypatch) + _patch_pusher_worker_drain(monkeypatch) + receive_bytes = AsyncMock( + side_effect=[ + _transcript_delivery_frame('delivery-1'), + pusher_router.WebSocketDisconnect(code=1000), + ] + ) + websocket = _PusherLifecycleWebSocket(receive_bytes=receive_bytes) + complete = MagicMock(return_value=True) + monkeypatch.setattr( + pusher_router.redis_db, + 'begin_pusher_delivery', + MagicMock(return_value=('claimed', 'lease-1')), + ) + monkeypatch.setattr(pusher_router.redis_db, 'complete_pusher_delivery', complete) + abandon = MagicMock(return_value=True) + monkeypatch.setattr(pusher_router.redis_db, 'abandon_pusher_delivery', abandon) + monkeypatch.setattr( + pusher_router, + 'trigger_realtime_integrations', + AsyncMock(side_effect=RuntimeError('delivery failed')), + ) + webhook = AsyncMock() + monkeypatch.setattr(pusher_router, 'realtime_transcript_webhook', webhook) + + await pusher_router._websocket_util_trigger(websocket, 'uid-1') + + complete.assert_not_called() + abandon.assert_called_once() + webhook.assert_not_awaited() + assert websocket.sent == [] + + +@pytest.mark.anyio +async def test_pusher_retryable_speaker_result_releases_lease_without_acknowledging(monkeypatch): + _PusherJourneyAttempt.outcomes = [] + _patch_pusher_session_dependencies(monkeypatch) + _patch_pusher_worker_drain(monkeypatch) + receive_bytes = AsyncMock( + side_effect=[ + _speaker_delivery_frame('speaker-delivery-1'), + struct.pack('I', pusher_router.PUSHER_DELIVERY_DRAIN_OPCODE), + pusher_router.WebSocketDisconnect(code=1000), + ] + ) + websocket = _PusherLifecycleWebSocket(receive_bytes=receive_bytes) + complete = MagicMock(return_value=True) + abandon = MagicMock(return_value=True) + monkeypatch.setattr( + pusher_router.redis_db, + 'begin_pusher_delivery', + MagicMock(return_value=('claimed', 'lease-1')), + ) + monkeypatch.setattr(pusher_router.redis_db, 'complete_pusher_delivery', complete) + monkeypatch.setattr(pusher_router.redis_db, 'abandon_pusher_delivery', abandon) + extract = AsyncMock(return_value=SpeakerSampleExtractionResult('retryable', 'audio_files_not_ready')) + monkeypatch.setattr(pusher_router, 'extract_speaker_samples', extract) + + await pusher_router._websocket_util_trigger(websocket, 'uid-1') + + extract.assert_awaited_once_with( + uid='uid-1', + person_id='person-1', + conversation_id='conversation-1', + segment_ids=['segment-1'], + sample_rate=8_000, + delivery_id='speaker-delivery-1', + ) + complete.assert_not_called() + abandon.assert_called_once() + assert websocket.sent == [] + + +@pytest.mark.anyio +async def test_pusher_drain_flushes_audio_metadata_before_speaker_extraction(monkeypatch): + _PusherJourneyAttempt.outcomes = [] + _patch_pusher_session_dependencies(monkeypatch) + _patch_pusher_worker_drain(monkeypatch) + monkeypatch.setattr(pusher_router.users_db, 'get_user_private_cloud_sync_enabled', lambda _uid: True) + monkeypatch.setattr(pusher_router.users_db, 'get_data_protection_level', lambda _uid: 'standard') + monkeypatch.setattr(pusher_router, 'is_audio_merge_dispatch_enabled', lambda: False) + upload_started = asyncio.Event() + release_upload = asyncio.Event() + extraction_started = asyncio.Event() + events: list[str] = [] + + def upload_audio(*_args, **_kwargs): + raise AssertionError('upload must run through the controlled executor boundary') + + monkeypatch.setattr(pusher_router, 'upload_audio_chunks_batch', upload_audio) + + async def controlled_run_blocking(_executor, func, *args, **kwargs): + if func is upload_audio: + events.append('upload_started') + upload_started.set() + await release_upload.wait() + events.append('upload_completed') + return None + return func(*args, **kwargs) + + monkeypatch.setattr(pusher_router, 'run_blocking', controlled_run_blocking) + monkeypatch.setattr( + pusher_router.conversations_db, + 'create_audio_files_from_chunks', + lambda _uid, _conversation_id: [SimpleNamespace(model_dump=lambda: {'id': 'audio-1'})], + ) + + def update_conversation(_uid, _conversation_id, payload): + assert payload == {'audio_files': [{'id': 'audio-1'}]} + events.append('metadata_updated') + + monkeypatch.setattr(pusher_router.conversations_db, 'update_conversation', update_conversation) + + async def extract(**_kwargs): + events.append('speaker_extraction') + extraction_started.set() + return SpeakerSampleExtractionResult('stored', 'sample_stored') + + monkeypatch.setattr(pusher_router, 'extract_speaker_samples', extract) + monkeypatch.setattr( + pusher_router.redis_db, + 'begin_pusher_delivery', + MagicMock(return_value=('claimed', 'lease-1')), + ) + monkeypatch.setattr(pusher_router.redis_db, 'complete_pusher_delivery', MagicMock(return_value=True)) + + receive_bytes = AsyncMock( + side_effect=[ + struct.pack('I', 103) + b'conversation-1', + struct.pack('I', 101) + struct.pack('d', 1000.0) + b'\x00\x01' * 16, + _speaker_delivery_frame('speaker-delivery-1'), + struct.pack('I', pusher_router.PUSHER_DELIVERY_DRAIN_OPCODE), + pusher_router.WebSocketDisconnect(code=1000), + ] + ) + websocket = _PusherLifecycleWebSocket(receive_bytes=receive_bytes) + + session_task = asyncio.create_task(pusher_router._websocket_util_trigger(websocket, 'uid-1')) + await asyncio.wait_for(upload_started.wait(), timeout=1) + await asyncio.sleep(0) + assert not extraction_started.is_set() + release_upload.set() + await asyncio.wait_for(session_task, timeout=2) + + assert events == ['upload_started', 'upload_completed', 'metadata_updated', 'speaker_extraction'] + acknowledgements = [json.loads(frame[4:]) for frame in websocket.sent if int.from_bytes(frame[:4], 'little') == 202] + assert acknowledgements == [{'kind': 'speaker_sample', 'delivery_id': 'speaker-delivery-1'}] + + +@pytest.mark.anyio +async def test_pusher_done_marker_suppresses_cross_socket_duplicate_and_reacks(monkeypatch): + _PusherJourneyAttempt.outcomes = [] + _patch_pusher_session_dependencies(monkeypatch) + _patch_pusher_worker_drain(monkeypatch) + receive_bytes = AsyncMock( + side_effect=[ + _transcript_delivery_frame('delivery-1'), + pusher_router.WebSocketDisconnect(code=1000), + ] + ) + websocket = _PusherLifecycleWebSocket(receive_bytes=receive_bytes) + monkeypatch.setattr( + pusher_router.redis_db, + 'begin_pusher_delivery', + MagicMock(return_value=('done', None)), + ) + complete = MagicMock() + integration = AsyncMock() + webhook = AsyncMock() + monkeypatch.setattr(pusher_router.redis_db, 'complete_pusher_delivery', complete) + monkeypatch.setattr(pusher_router, 'trigger_realtime_integrations', integration) + monkeypatch.setattr(pusher_router, 'realtime_transcript_webhook', webhook) + + await pusher_router._websocket_util_trigger(websocket, 'uid-1') + + integration.assert_not_awaited() + webhook.assert_not_awaited() + complete.assert_not_called() + assert len(websocket.sent) == 1 + assert int.from_bytes(websocket.sent[0][:4], 'little') == 202 + + +@pytest.mark.anyio +async def test_pusher_full_stable_queue_rejects_without_evicting_or_acknowledging(monkeypatch): + _PusherJourneyAttempt.outcomes = [] + _patch_pusher_session_dependencies(monkeypatch) + _patch_pusher_worker_drain(monkeypatch) + monkeypatch.setattr(pusher_router, 'TRANSCRIPT_QUEUE_WARN_SIZE', 1) + receive_bytes = AsyncMock( + side_effect=[ + _transcript_delivery_frame('delivery-1', segment_id='first'), + _transcript_delivery_frame('delivery-2', segment_id='second'), + pusher_router.WebSocketDisconnect(code=1000), + ] + ) + websocket = _PusherLifecycleWebSocket(receive_bytes=receive_bytes) + begin = MagicMock(return_value=('claimed', 'lease-1')) + integration = AsyncMock() + monkeypatch.setattr(pusher_router.redis_db, 'begin_pusher_delivery', begin) + monkeypatch.setattr(pusher_router.redis_db, 'complete_pusher_delivery', MagicMock(return_value=True)) + monkeypatch.setattr(pusher_router, 'trigger_realtime_integrations', integration) + monkeypatch.setattr(pusher_router, 'realtime_transcript_webhook', AsyncMock()) + + await pusher_router._websocket_util_trigger(websocket, 'uid-1') + + assert begin.call_count == 1 + assert integration.await_count == 1 + assert integration.await_args.args[1] == [{'id': 'first', 'text': 'hello'}] + acknowledgements = [json.loads(frame[4:]) for frame in websocket.sent if int.from_bytes(frame[:4], 'little') == 202] + assert acknowledgements == [{'kind': 'transcript', 'delivery_id': 'delivery-1'}] + + +@pytest.mark.anyio +async def test_pusher_legacy_overflow_cannot_evict_an_accepted_stable_delivery(monkeypatch): + _PusherJourneyAttempt.outcomes = [] + _patch_pusher_session_dependencies(monkeypatch) + _patch_pusher_worker_drain(monkeypatch) + monkeypatch.setattr(pusher_router, 'TRANSCRIPT_QUEUE_WARN_SIZE', 1) + legacy_payload = { + 'segments': [{'id': 'legacy', 'text': 'legacy'}], + 'memory_id': 'conversation-1', + } + receive_bytes = AsyncMock( + side_effect=[ + _transcript_delivery_frame('delivery-1', segment_id='stable'), + struct.pack('I', 102) + json.dumps(legacy_payload).encode('utf-8'), + pusher_router.WebSocketDisconnect(code=1000), + ] + ) + websocket = _PusherLifecycleWebSocket(receive_bytes=receive_bytes) + begin = MagicMock(return_value=('claimed', 'lease-1')) + integration = AsyncMock() + monkeypatch.setattr(pusher_router.redis_db, 'begin_pusher_delivery', begin) + monkeypatch.setattr(pusher_router.redis_db, 'complete_pusher_delivery', MagicMock(return_value=True)) + monkeypatch.setattr(pusher_router, 'trigger_realtime_integrations', integration) + monkeypatch.setattr(pusher_router, 'realtime_transcript_webhook', AsyncMock()) + + await pusher_router._websocket_util_trigger(websocket, 'uid-1') + + begin.assert_called_once() + integration.assert_awaited_once() + assert integration.await_args.args[1] == [{'id': 'stable', 'text': 'hello'}] + acknowledgements = [json.loads(frame[4:]) for frame in websocket.sent if int.from_bytes(frame[:4], 'little') == 202] + assert acknowledgements == [{'kind': 'transcript', 'delivery_id': 'delivery-1'}] + + @pytest.mark.anyio async def test_pusher_claims_the_durable_job_before_finalizing(monkeypatch): websocket = _PusherWebSocket() @@ -802,7 +1134,7 @@ async def test_pusher_keeps_a_completed_job_terminal_when_source_result_delivery websocket = _PusherWebSocket() async def closed_send(_payload: bytes) -> None: - raise RuntimeError('Cannot call send once closed') + raise OSError('socket closed before result delivery') websocket.send_bytes = closed_send completed = MagicMock(return_value=True) diff --git a/backend/tests/unit/test_lock_bypass_fixes.py b/backend/tests/unit/test_lock_bypass_fixes.py index 1652ab35c1d..ecdf1f47098 100644 --- a/backend/tests/unit/test_lock_bypass_fixes.py +++ b/backend/tests/unit/test_lock_bypass_fixes.py @@ -1639,7 +1639,7 @@ def test_mcp_delete_memory_allows_unlocked(self): import database.memories as memories_db memories_db.get_memory = MagicMock(return_value=_make_memory(locked=False)) - memories_db.delete_memory = MagicMock() + memories_db.delete_memory = MagicMock(return_value=SimpleNamespace(committed_count=1)) from routers import mcp from routers.mcp import delete_memory @@ -1729,16 +1729,25 @@ def test_mcp_sse_delete_memory_allows_unlocked(self): import database.memories as memories_db memories_db.get_memory = MagicMock(return_value=_make_memory(locked=False)) - memories_db.delete_memory = MagicMock() + memories_db.delete_memory = MagicMock(return_value=SimpleNamespace(committed_count=1)) from routers import mcp_sse from routers.mcp_sse import execute_tool _allow_memory_product_auth(mcp_sse) _force_legacy_memory_paths(mcp_sse) - result = execute_tool('test-uid', 'delete_memory', {'memory_id': 'mem-1'}, auth_context=_memory_auth_context()) + result = execute_tool( + 'test-uid', + 'delete_memory', + {'memory_id': 'mem-1'}, + auth_context=_memory_auth_context(), + ) assert result == {"success": True} - memories_db.delete_memory.assert_called_once_with('test-uid', 'mem-1') + memories_db.delete_memory.assert_called_once_with( + 'test-uid', + 'mem-1', + firestore_client=mcp_sse.db, + ) def test_mcp_sse_delete_memory_404_missing(self): import database.memories as memories_db @@ -1789,7 +1798,12 @@ def test_mcp_sse_edit_memory_allows_unlocked(self): 'test-uid', 'edit_memory', {'memory_id': 'mem-1', 'content': 'new'}, auth_context=_memory_auth_context() ) assert result == {"success": True} - memories_db.edit_memory.assert_called_once_with('test-uid', 'mem-1', 'new') + memories_db.edit_memory.assert_called_once_with( + 'test-uid', + 'mem-1', + 'new', + firestore_client=mcp_sse.db, + ) # ============================================================================= diff --git a/backend/tests/unit/test_mcp_search_memories.py b/backend/tests/unit/test_mcp_search_memories.py index c026c2653ab..e2f6f0ef544 100644 --- a/backend/tests/unit/test_mcp_search_memories.py +++ b/backend/tests/unit/test_mcp_search_memories.py @@ -4,7 +4,7 @@ following the pattern in test_lock_bypass_fixes.py. """ -from unittest.mock import patch, MagicMock +from unittest.mock import ANY, MagicMock, patch import os import pytest import sys @@ -408,7 +408,9 @@ def test_edit_upserts_vector_with_new_content( mock_service_memories_db.get_memory.return_value = _legacy_memory_doc(category='hobbies') result = edit_memory(memory_id="mem-1", value="new text", auth_context=_auth_context()) assert result == {"status": "ok"} - mock_service_memories_db.edit_memory.assert_called_once_with("user-1", "mem-1", "new text") + mock_service_memories_db.edit_memory.assert_called_once_with( + "user-1", "mem-1", "new text", firestore_client=ANY + ) mock_upsert_vector.assert_called_once_with("user-1", "mem-1", "new text", "hobbies", subject_entity_id=None) @patch('utils.memory.memory_service.guard_legacy_memory_write', side_effect=_allowed_write_guard) @@ -425,7 +427,9 @@ def test_edit_succeeds_when_vector_upsert_fails( mock_upsert_vector.side_effect = Exception("pinecone down") result = edit_memory(memory_id="mem-1", value="new text", auth_context=_auth_context()) assert result == {"status": "ok"} - mock_service_memories_db.edit_memory.assert_called_once_with("user-1", "mem-1", "new text") + mock_service_memories_db.edit_memory.assert_called_once_with( + "user-1", "mem-1", "new text", firestore_client=ANY + ) class TestDeleteMemoryVectorSync: @@ -448,9 +452,12 @@ def test_delete_removes_vector( mock_pin.return_value = _LEGACY mock_fetch.return_value = {'id': 'mem-1', 'content': 'x', 'is_locked': False} mock_service_memories_db.get_memory.return_value = {'id': 'mem-1', 'content': 'x', 'is_locked': False} + mock_service_memories_db.delete_memory.return_value = SimpleNamespace( + committed_count=1, + ) result = delete_memory(memory_id="mem-1", auth_context=_auth_context()) assert result == {"status": "ok"} - mock_service_memories_db.delete_memory.assert_called_once_with("user-1", "mem-1") + mock_service_memories_db.delete_memory.assert_called_once_with("user-1", "mem-1", firestore_client=ANY) mock_delete_vector.assert_called_once_with("user-1", "mem-1") @patch('utils.memory.memory_service.guard_legacy_memory_write', side_effect=_allowed_write_guard) @@ -464,7 +471,10 @@ def test_delete_succeeds_when_vector_delete_fails( mock_pin.return_value = _LEGACY mock_fetch.return_value = {'id': 'mem-1', 'content': 'x', 'is_locked': False} mock_service_memories_db.get_memory.return_value = {'id': 'mem-1', 'content': 'x', 'is_locked': False} + mock_service_memories_db.delete_memory.return_value = SimpleNamespace( + committed_count=1, + ) mock_delete_vector.side_effect = Exception("pinecone down") result = delete_memory(memory_id="mem-1", auth_context=_auth_context()) assert result == {"status": "ok"} - mock_service_memories_db.delete_memory.assert_called_once_with("user-1", "mem-1") + mock_service_memories_db.delete_memory.assert_called_once_with("user-1", "mem-1", firestore_client=ANY) diff --git a/backend/tests/unit/test_memories_batch.py b/backend/tests/unit/test_memories_batch.py index cfdab1163e8..d468f422ce4 100644 --- a/backend/tests/unit/test_memories_batch.py +++ b/backend/tests/unit/test_memories_batch.py @@ -14,58 +14,25 @@ import pytest -# Stub heavy deps before importing vector_db / routers.memories. These -# modules pull in `pinecone`, `google.cloud.firestore`, `firebase_admin`, and -# `utils.llm.clients.embeddings` at import time, none of which are available -# (or desirable) in the unit test environment. +# Stub only the dependencies imported directly by vector_db. Firestore and +# firebase are not part of this module's import path, and replacing their +# parent packages here prevents later test modules from importing real +# google.cloud subpackages during collection. for mod_name in [ 'pinecone', - 'firebase_admin', - 'firebase_admin.auth', - 'google', - 'google.cloud', - 'google.cloud.firestore', ]: if mod_name not in sys.modules: sys.modules[mod_name] = types.ModuleType(mod_name) sys.modules['pinecone'].Pinecone = MagicMock - -class _FakeFirestoreClient: - def collection(self, *a, **kw): - return MagicMock() - - def batch(self): - return MagicMock() - - -sys.modules['google.cloud.firestore'].Client = _FakeFirestoreClient -sys.modules['google.cloud.firestore'].ArrayUnion = MagicMock -sys.modules['google.cloud.firestore'].ArrayRemove = MagicMock -sys.modules['google.cloud.firestore'].Increment = MagicMock -sys.modules['google.cloud.firestore'].SERVER_TIMESTAMP = object() -sys.modules['google.cloud.firestore'].DELETE_FIELD = object() -sys.modules['google.cloud.firestore'].FieldFilter = MagicMock -sys.modules['google.cloud.firestore'].Query = MagicMock -sys.modules['firebase_admin.auth'].InvalidIdTokenError = type('InvalidIdTokenError', (Exception,), {}) - -# Stub `utils.llm.clients.embeddings` only. Don't overwrite `utils` or -# `utils.llm` as packages — other real submodules (utils.rate_limit_config, -# utils.other.endpoints) must remain importable. -if 'utils.llm.clients' not in sys.modules: - clients_stub = types.ModuleType('utils.llm.clients') - clients_stub.embeddings = MagicMock() - sys.modules['utils.llm.clients'] = clients_stub - - from database import vector_db # noqa: E402 class TestUpsertMemoryVectorsBatch: def _setup_mocks(self, monkeypatch, *, index_none=False): fake_index = MagicMock() - fake_index.upsert = MagicMock(return_value={'upserted_count': 2}) + fake_index.upsert = MagicMock(side_effect=lambda *, vectors, namespace: {'upserted_count': len(vectors)}) monkeypatch.setattr(vector_db, 'index', None if index_none else fake_index) fake_embeddings = MagicMock() @@ -106,6 +73,16 @@ def test_single_upsert_strips_null_projection_metadata(self, monkeypatch): assert metadata['projection_version'] == 1 assert metadata['source_tombstone_state'] == 'active' + def test_single_upsert_returns_nonwrite_when_pinecone_reports_zero(self, monkeypatch): + fake_index, fake_embeddings = self._setup_mocks(monkeypatch) + fake_embeddings.embed_query = MagicMock(return_value=[0.1, 0.2]) + fake_index.upsert.side_effect = None + fake_index.upsert.return_value = {'upserted_count': 0} + + written = vector_db.upsert_memory_vector('uid-abc', 'm1', 'hello', 'system') + + assert written is None + def test_batch_upsert_uses_single_embed_and_single_upsert(self, monkeypatch): """The whole point of the helper: one embed call + one upsert call.""" fake_index, fake_embeddings = self._setup_mocks(monkeypatch) @@ -159,6 +136,21 @@ def test_batch_upsert_strips_null_projection_metadata(self, monkeypatch): assert 'valid_time' not in metadata assert metadata['projection_version'] == 1 + def test_batch_upsert_returns_reported_partial_write_count(self, monkeypatch): + fake_index, _ = self._setup_mocks(monkeypatch) + fake_index.upsert.side_effect = None + fake_index.upsert.return_value = {'upserted_count': 1} + + written = vector_db.upsert_memory_vectors_batch( + 'uid-abc', + [ + {'memory_id': 'm1', 'content': 'apple', 'category': 'manual'}, + {'memory_id': 'm2', 'content': 'banana', 'category': 'manual'}, + ], + ) + + assert written == 1 + def test_batch_upsert_empty_list_is_noop(self, monkeypatch): fake_index, fake_embeddings = self._setup_mocks(monkeypatch) diff --git a/backend/tests/unit/test_memories_batch_delete.py b/backend/tests/unit/test_memories_batch_delete.py index 9bb0391489f..ef86d04ef2c 100644 --- a/backend/tests/unit/test_memories_batch_delete.py +++ b/backend/tests/unit/test_memories_batch_delete.py @@ -42,10 +42,14 @@ def _force_legacy(monkeypatch): def _patch_db(monkeypatch, fetched): get_mock = MagicMock(return_value=fetched) - delete_mock = MagicMock() + delete_mock = MagicMock( + return_value=mem_mod.memories_db.LegacyMemoryDeleteResult( + memory_ids=[item["id"] for item in fetched], + ) + ) monkeypatch.setattr(mem_mod.memories_db, 'get_memories_by_ids', get_mock) monkeypatch.setattr(mem_mod.memories_db, 'delete_memories_batch', delete_mock) - vectors_mock = MagicMock() + vectors_mock = MagicMock(return_value=len(fetched)) monkeypatch.setattr(mem_mod, 'delete_memory_vectors_batch', vectors_mock) return get_mock, delete_mock, vectors_mock @@ -189,6 +193,27 @@ def test_vector_delete_failure_does_not_fail_the_request(self, monkeypatch): delete_mock.assert_called_once_with('u1', ['a']) +class TestLegacyDeleteAllProjection: + def test_delete_all_uses_ids_from_the_atomic_delete_snapshot(self, monkeypatch): + _force_legacy(monkeypatch) + all_ids = ["visible", "locked", "rejected", "invalidated"] + delete_all = MagicMock( + return_value=mem_mod.memories_db.LegacyMemoryDeleteResult( + memory_ids=all_ids, + ) + ) + monkeypatch.setattr(mem_mod.memories_db, "delete_all_memories", delete_all) + monkeypatch.setattr( + mem_mod.memories_db, + "get_memories", + MagicMock(side_effect=AssertionError("delete-all must not use the filtered memory reader")), + ) + monkeypatch.setattr(mem_mod, "delete_memory_vectors_batch", MagicMock(return_value=2)) + + assert mem_mod.delete_memories(uid="u1") == {"status": "ok"} + delete_all.assert_called_once_with("u1") + + class TestBatchDeleteCanonicalCohort: def test_canonical_cohort_mirrors_single_delete_canonical_path(self, monkeypatch): # Canonical cohort delegates the full selection to one atomic adapter call and diff --git a/backend/tests/unit/test_memories_delete_batch_chunk.py b/backend/tests/unit/test_memories_delete_batch_chunk.py index 876d36a62e4..871a91126c5 100644 --- a/backend/tests/unit/test_memories_delete_batch_chunk.py +++ b/backend/tests/unit/test_memories_delete_batch_chunk.py @@ -3,9 +3,8 @@ database.memories.delete_memories and delete_all_memories accumulated every delete into a single WriteBatch and committed once. Firestore rejects a batch with more than 500 writes, so a user with more than 500 memories made batch.commit() raise, and the delete (including the -account-deletion path) removed nothing. Both functions now chunk at 499, mirroring -unlock_all_memories. The fake below models the real 500-write limit by raising on an oversized -commit, so the pre-fix single-batch code fails here. +account-deletion path) removed nothing. The fake below models the real 500-write limit by +raising on an oversized commit, so the pre-fix single-batch code fails here. """ import database.memories as memories @@ -14,17 +13,23 @@ class _FakeBatch: - def __init__(self, commit_sink): + def __init__(self, commit_sink, record_sink): self._commit_sink = commit_sink + self._record_sink = record_sink self.deletes = 0 + self.sets = 0 def delete(self, reference): self.deletes += 1 + def set(self, reference, payload): + self.sets += 1 + self._record_sink.append(dict(payload)) + def commit(self): - if self.deletes > _FIRESTORE_BATCH_LIMIT: + if self.deletes + self.sets > _FIRESTORE_BATCH_LIMIT: raise ValueError("Firestore batch too large: max 500 writes per commit") - self._commit_sink.append(self.deletes) + self._commit_sink.append((self.deletes, self.sets)) class _FakeDoc: @@ -51,23 +56,26 @@ class _FakeDb: def __init__(self, n_docs, commit_sink): self._docs = [_FakeDoc(i) for i in range(n_docs)] self._commit_sink = commit_sink + self.staged_records = [] def collection(self, _name): return _FakeCollection(self._docs) def batch(self): - return _FakeBatch(self._commit_sink) + return _FakeBatch(self._commit_sink, self.staged_records) def test_delete_all_memories_chunks_over_firestore_batch_limit(): commit_sink = [] fake = _FakeDb(1000, commit_sink) - memories.delete_all_memories("u1", firestore_client=fake) # must not raise + result = memories.delete_all_memories("u1", firestore_client=fake) # must not raise - assert sum(commit_sink) == 1000 # every memory deleted + assert sum(deletes for deletes, _sets in commit_sink) == 1000 # every memory deleted + assert sum(sets for _deletes, sets in commit_sink) == 0 + assert result.committed_count == 1000 assert len(commit_sink) >= 2 # split across batches - assert all(c <= _FIRESTORE_BATCH_LIMIT for c in commit_sink) + assert all(deletes + sets <= _FIRESTORE_BATCH_LIMIT for deletes, sets in commit_sink) def test_delete_memories_chunks_over_firestore_batch_limit(): @@ -76,15 +84,31 @@ def test_delete_memories_chunks_over_firestore_batch_limit(): memories.delete_memories("u1", firestore_client=fake) # must not raise - assert sum(commit_sink) == 1000 + assert sum(deletes for deletes, _sets in commit_sink) == 1000 + assert sum(sets for _deletes, sets in commit_sink) == 0 assert len(commit_sink) >= 2 - assert all(c <= _FIRESTORE_BATCH_LIMIT for c in commit_sink) + assert all(deletes + sets <= _FIRESTORE_BATCH_LIMIT for deletes, sets in commit_sink) + + +def test_delete_memories_batch_returns_verified_delete_count(): + commit_sink = [] + fake = _FakeDb(0, commit_sink) + + result = memories.delete_memories_batch( + "u1", + [f"mem-{index}" for index in range(1000)], + firestore_client=fake, + ) + + assert result.committed_count == 1000 + assert commit_sink == [(499, 0), (499, 0), (2, 0)] def test_delete_all_memories_small_count_single_commit(): commit_sink = [] fake = _FakeDb(3, commit_sink) - memories.delete_all_memories("u1", firestore_client=fake) + result = memories.delete_all_memories("u1", firestore_client=fake) - assert commit_sink == [3] + assert commit_sink == [(3, 0)] + assert result.committed_count == 3 diff --git a/backend/tests/unit/test_memory_ledger.py b/backend/tests/unit/test_memory_ledger.py index 016543a9633..653606cd2ac 100644 --- a/backend/tests/unit/test_memory_ledger.py +++ b/backend/tests/unit/test_memory_ledger.py @@ -29,8 +29,7 @@ @pytest.fixture(scope="module", autouse=True) def _load_modules(): - """Load fresh database.memory_ledger (+ its database.projection_repair dep) against - a stubbed database._client + google.cloud.firestore_v1 chain.""" + """Load the ledger and legacy projection helpers against stubbed Firestore.""" client_stub = ModuleType("database._client") client_stub.db = MagicMock(name="db") client_stub.document_id_from_seed = lambda seed: "id-" + str(abs(hash(seed)) % (10**12)) @@ -49,11 +48,14 @@ def _load_modules(): "google.cloud.firestore_v1": firestore_v1_stub, } with stub_modules(fakes): + projection_repair = load_module_fresh( + "database.projection_repair", + os.path.join(str(_BACKEND), "database", "projection_repair.py"), + ) memory_ledger = load_module_fresh( "database.memory_ledger", os.path.join(str(_BACKEND), "database", "memory_ledger.py"), ) - projection_repair = memory_ledger.projection_repair globals()["memory_ledger"] = memory_ledger globals()["projection_repair"] = projection_repair yield @@ -814,6 +816,55 @@ def fast_retry(transaction_factory, operation, **kwargs): assert queued == [] +def test_enqueue_projection_repairs_chunks_large_commit(): + commit_sizes = [] + + class FakeSnapshot: + exists = False + + class FakeDocument: + def collection(self, _value): + return self + + def document(self, _value): + return self + + def get(self): + return FakeSnapshot() + + class FakeBatch: + def __init__(self): + self.write_count = 0 + + def set(self, _ref, _payload): + self.write_count += 1 + + def commit(self): + if self.write_count > 500: + raise ValueError("Firestore batch too large") + commit_sizes.append(self.write_count) + + class FakeDB: + def collection(self, _value): + return FakeDocument() + + def batch(self): + return FakeBatch() + + memory_ids = [f"memory-{index}" for index in range(1000)] + repair_ids = projection_repair.enqueue_projection_repairs( + "uid-1", + { + "commit_id": "large-commit", + "mutations": [{"type": "retract_fact", "fact_id": memory_id} for memory_id in memory_ids], + }, + firestore_client=FakeDB(), + ) + + assert len(repair_ids) == 1000 + assert commit_sizes == [499, 499, 2] + + def test_process_projection_repairs_applies_queued_vector_repairs(monkeypatch): updates = [] diff --git a/backend/tests/unit/test_memory_service_parity.py b/backend/tests/unit/test_memory_service_parity.py index 41de1f5ced4..6a4e5f0a01d 100644 --- a/backend/tests/unit/test_memory_service_parity.py +++ b/backend/tests/unit/test_memory_service_parity.py @@ -260,11 +260,88 @@ def test_search_matches_direct_legacy_helper(self, monkeypatch): direct = service_mod._legacy_search_memories("uid-test", "hiking", limit=5) assert find_similar.call_count == 2 - find_similar.assert_called_with("uid-test", "hiking", threshold=0.0, limit=5) + find_similar.assert_called_with("uid-test", "hiking", threshold=0.0, limit=15) assert get_by_ids.call_count == 2 get_by_ids.assert_called_with("uid-test", ["mem-1", "mem-2"]) assert via_service == direct + def test_search_overfetches_and_backfills_filtered_or_invalid_vector_hits(self, monkeypatch): + service_mod = _load_memory_service(monkeypatch) + vector_matches = [ + {"memory_id": "missing", "score": 0.99}, + {"memory_id": "locked", "score": 0.98}, + {"memory_id": "rejected", "score": 0.97}, + {"memory_id": "invalidated", "score": 0.96}, + {"memory_id": "malformed", "score": 0.95}, + {"memory_id": "visible-1", "score": 0.80}, + {"memory_id": "visible-2", "score": 0.70}, + ] + locked = _sample_memory_dict("locked", locked=True) + rejected = _sample_memory_dict("rejected") + rejected["user_review"] = False + invalidated = _sample_memory_dict("invalidated") + invalidated["invalid_at"] = datetime(2026, 1, 16, tzinfo=timezone.utc) + malformed = _sample_memory_dict("malformed") + malformed.pop("content") + # Firestore get_all order is not a relevance contract. Deliberately + # return visible docs in reverse order and require vector ordering. + memories = [ + _sample_memory_dict("visible-2"), + invalidated, + rejected, + malformed, + locked, + _sample_memory_dict("visible-1"), + ] + + def find_similar_memories(_uid, _query, *, threshold, limit): + assert threshold == 0.0 + return vector_matches[:limit] + + def get_memories_by_ids(_uid, memory_ids): + requested = set(memory_ids) + return [memory for memory in memories if memory["id"] in requested] + + with ( + patch.object( + service_mod.vector_db, + "find_similar_memories", + side_effect=find_similar_memories, + ) as find_similar, + patch.object(service_mod.memories_db, "get_memories_by_ids", side_effect=get_memories_by_ids), + ): + results = service_mod._legacy_search_memories("uid-test", "hiking", limit=2) + + assert [call.kwargs["limit"] for call in find_similar.call_args_list] == [ + 6, + 12, + ], "memory search must overfetch until filtering yields requested visible results" + assert [result.memory.id for result in results] == ["visible-1", "visible-2"] + assert [result.score for result in results] == [0.80, 0.70] + + def test_search_progressive_backfill_stops_at_candidate_cap(self, monkeypatch): + service_mod = _load_memory_service(monkeypatch) + + def find_similar_memories(_uid, _query, *, threshold, limit): + assert threshold == 0.0 + return [{"memory_id": f"locked-{index}", "score": 1.0 - (index / 100)} for index in range(limit)] + + def get_memories_by_ids(_uid, memory_ids): + return [_sample_memory_dict(memory_id, locked=True) for memory_id in memory_ids] + + with ( + patch.object( + service_mod.vector_db, + "find_similar_memories", + side_effect=find_similar_memories, + ) as find_similar, + patch.object(service_mod.memories_db, "get_memories_by_ids", side_effect=get_memories_by_ids), + ): + results = service_mod._legacy_search_memories("uid-test", "hiking", limit=2) + + assert results == [] + assert [call.kwargs["limit"] for call in find_similar.call_args_list] == [6, 12, 24, 48, 60] + def test_external_canonical_write_gate_failure_does_not_fallback_to_legacy_create(self, monkeypatch): service_mod = _load_memory_service(monkeypatch) memory_db = service_mod.MemoryDB.model_validate(_sample_memory_dict()) @@ -512,7 +589,8 @@ def test_external_legacy_create_strips_canonical_lifecycle_fields(self, monkeypa lambda *args, **kwargs: SimpleNamespace(allowed=True, status_code=200, detail=None), ) - result = service_mod.MemoryService(db_client=_FirestoreFake()).create_external_memory( + db_client = _FirestoreFake() + result = service_mod.MemoryService(db_client=db_client).create_external_memory( "uid-test", memory_db, memory_system=MemorySystem.LEGACY, @@ -526,6 +604,80 @@ def test_external_legacy_create_strips_canonical_lifecycle_fields(self, monkeypa assert "layer" not in payload assert "tier" not in payload assert result.memory_tier is None + assert create_memory.call_args.kwargs["firestore_client"] is db_client + + def test_external_legacy_edit_uses_the_injected_firestore_client_for_reads_and_write(self, monkeypatch): + service_mod = _load_memory_service(monkeypatch) + before = _sample_memory_dict() + after = {**before, "content": "Updated hiking preference"} + monkeypatch.setattr( + service_mod, + "guard_legacy_memory_write", + lambda *args, **kwargs: SimpleNamespace(allowed=True, status_code=200, detail=None), + ) + get_memory = MagicMock(side_effect=[before, after]) + monkeypatch.setattr(service_mod.memories_db, "get_memory", get_memory) + monkeypatch.setattr( + service_mod.memories_db, + "edit_memory", + MagicMock(return_value={"commit": {"commit_id": "commit-edit"}}), + ) + db_client = _FirestoreFake() + + result = service_mod.MemoryService(db_client=db_client).update_external_memory_content( + "uid-test", + "mem-1", + "Updated hiking preference", + memory_system=MemorySystem.LEGACY, + consumer="mcp", + operation="mcp_tool_memory_edit", + upsert_vector=False, + ) + + assert result.content == "Updated hiking preference" + assert [item.kwargs["firestore_client"] for item in get_memory.call_args_list] == [db_client, db_client] + service_mod.memories_db.edit_memory.assert_called_once_with( + "uid-test", + "mem-1", + "Updated hiking preference", + firestore_client=db_client, + ) + + def test_external_legacy_delete_uses_the_injected_firestore_client_for_read_and_write(self, monkeypatch): + service_mod = _load_memory_service(monkeypatch) + monkeypatch.setattr( + service_mod, + "guard_legacy_memory_write", + lambda *args, **kwargs: SimpleNamespace(allowed=True, status_code=200, detail=None), + ) + get_memory = MagicMock(return_value=_sample_memory_dict()) + monkeypatch.setattr(service_mod.memories_db, "get_memory", get_memory) + monkeypatch.setattr( + service_mod.memories_db, + "delete_memory", + MagicMock( + return_value=service_mod.memories_db.LegacyMemoryDeleteResult( + memory_ids=["mem-1"], + ) + ), + ) + db_client = _FirestoreFake() + + service_mod.MemoryService(db_client=db_client).delete_external_memory( + "uid-test", + "mem-1", + memory_system=MemorySystem.LEGACY, + consumer="mcp", + operation="mcp_tool_memory_delete", + delete_vector=False, + ) + + get_memory.assert_called_once_with("uid-test", "mem-1", firestore_client=db_client) + service_mod.memories_db.delete_memory.assert_called_once_with( + "uid-test", + "mem-1", + firestore_client=db_client, + ) def test_search_mcp_legacy_fetch_limit_filters_and_rrf(self, monkeypatch): service_mod = _load_memory_service(monkeypatch) @@ -578,6 +730,43 @@ def test_search_mcp_legacy_fetch_limit_filters_and_rrf(self, monkeypatch): assert len(direct) == 1 assert direct[0]["id"] == "mem-ok" + def test_search_mcp_progressively_backfills_filtered_vector_hits(self, monkeypatch): + service_mod = _load_memory_service(monkeypatch) + vector_matches = [{"memory_id": f"rejected-{index}", "score": 0.99 - (index / 100)} for index in range(5)] + [ + {"memory_id": "visible-1", "score": 0.80}, + {"memory_id": "visible-2", "score": 0.70}, + ] + memories = [] + for index in range(5): + rejected = _sample_memory_dict(f"rejected-{index}") + rejected["user_review"] = False + memories.append(rejected) + memories.extend([_sample_memory_dict("visible-2"), _sample_memory_dict("visible-1")]) + + def find_similar_memories(_uid, _query, *, threshold, limit): + assert threshold == 0.0 + return vector_matches[:limit] + + def get_memories_by_ids(_uid, memory_ids): + requested = set(memory_ids) + return [memory for memory in memories if memory["id"] in requested] + + with ( + patch.object( + service_mod.vector_db, + "find_similar_memories", + side_effect=find_similar_memories, + ) as find_similar, + patch.object(service_mod.memories_db, "get_memories_by_ids", side_effect=get_memories_by_ids), + patch.object( + service_mod, "rrf_rerank", side_effect=lambda query, candidates, limit, k=60: candidates[:limit] + ), + ): + results = service_mod._legacy_search_memories_mcp("uid-test", "visible", limit=2) + + assert [call.kwargs["limit"] for call in find_similar.call_args_list] == [6, 12] + assert [result["id"] for result in results] == ["visible-1", "visible-2"] + class TestMemoryServiceUsesRequestPin: def test_search_mcp_stays_on_pinned_legacy_backend_when_resolver_flips(self, monkeypatch): diff --git a/backend/tests/unit/test_pusher_ghost_connections.py b/backend/tests/unit/test_pusher_ghost_connections.py index 214886215c4..ec6543823cb 100644 --- a/backend/tests/unit/test_pusher_ghost_connections.py +++ b/backend/tests/unit/test_pusher_ghost_connections.py @@ -822,13 +822,30 @@ def test_bg_main_tasks_has_four_tasks(self): def test_is_shutdown_guards_speaker_sample_age_check(self): """In process_speaker_sample_queue, is_shutdown must be checked in the same conditional as SPEAKER_SAMPLE_MIN_AGE.""" - src = _read_source(PUSHER_SRC) - lines = src.split('\n') - - for line in lines: - if 'is_shutdown' in line and 'SPEAKER_SAMPLE_MIN_AGE' in line: - return - pytest.fail("is_shutdown must guard the SPEAKER_SAMPLE_MIN_AGE check in process_speaker_sample_queue") + handler = _parse_handler_ast() + age_conditions = [ + node.test + for node in ast.walk(handler) + if isinstance(node, ast.If) + and any( + isinstance(child, ast.Name) and child.id == 'SPEAKER_SAMPLE_MIN_AGE' for child in ast.walk(node.test) + ) + ] + + assert len(age_conditions) == 1, "expected one speaker-sample age condition" + condition = age_conditions[0] + assert isinstance(condition, ast.BoolOp) and isinstance( + condition.op, ast.Or + ), "speaker-sample age condition must be bypassed with OR guards" + + terms = list(condition.values) + for term in list(terms): + if isinstance(term, ast.BoolOp) and isinstance(term.op, ast.Or): + terms.extend(term.values) + + assert any( + isinstance(term, ast.Name) and term.id == 'is_shutdown' for term in terms + ), "is_shutdown must bypass SPEAKER_SAMPLE_MIN_AGE in process_speaker_sample_queue" def test_drain_tasks_handles_timeout_logging(self): """drain_tasks utility handles timeout logging — verify it's used in pusher.""" diff --git a/backend/tests/unit/test_pusher_heartbeat.py b/backend/tests/unit/test_pusher_heartbeat.py index 2a50a3f36d8..41f937b2ad5 100644 --- a/backend/tests/unit/test_pusher_heartbeat.py +++ b/backend/tests/unit/test_pusher_heartbeat.py @@ -11,10 +11,6 @@ import pytest -from tests.unit.pusher_websockets_stub import install_websockets_stub - -install_websockets_stub() - from websockets.exceptions import ConnectionClosed pytestmark = pytest.mark.slow @@ -48,7 +44,13 @@ async def _simulate_pusher_dispatch(frames: list[bytes]) -> dict: This mirrors the real dispatch at backend/routers/pusher.py:324-451. """ - counts = {"heartbeat": 0, "conversation_id": 0, "transcript": 0, "audio": 0, "unknown": 0} + counts = { + "heartbeat": 0, + "conversation_id": 0, + "transcript": 0, + "audio": 0, + "unknown": 0, + } for data in frames: if len(data) < 4: @@ -293,6 +295,6 @@ def test_heartbeat_frame_is_minimal(): def test_heartbeat_header_does_not_collide_with_existing_headers(): """Header 100 is distinct from all existing protocol headers.""" - existing_headers = {101, 102, 103, 104, 105, 201} # All existing headers + existing_headers = {101, 102, 103, 104, 105, 201, 202} # All existing headers heartbeat_header = 100 assert heartbeat_header not in existing_headers, "Header 100 must not collide with existing protocol headers" diff --git a/backend/tests/unit/test_pusher_readiness_drain.py b/backend/tests/unit/test_pusher_readiness_drain.py index fa6b0571789..ba67747b337 100644 --- a/backend/tests/unit/test_pusher_readiness_drain.py +++ b/backend/tests/unit/test_pusher_readiness_drain.py @@ -227,7 +227,7 @@ def test_real_app_registers_readiness_routes_and_shutdown_drain(handlers): def test_ws_drain_reject_ordering_is_static_tripwire(): """STATIC source-ordering tripwire (NOT behavioral coverage): asserts that, in - the source of ``_websocket_util_trigger``, ``await websocket.accept()`` precedes + the source of ``_websocket_util_trigger``, ``await websocket.accept(...)`` precedes the ``ReadinessGate.is_serving()`` drain check, which precedes ``await websocket.close(code=1001)``. It does NOT drive a live WebSocket handshake or observe a real close frame; per AGENTS.md this is a static @@ -244,7 +244,7 @@ def test_ws_drain_reject_ordering_is_static_tripwire(): from routers import pusher as pusher_router source = inspect.getsource(pusher_router._websocket_util_trigger) - accept_pos = source.find('await websocket.accept()') + accept_pos = source.find('await websocket.accept(') drain_check_pos = source.find('if not ReadinessGate.is_serving()') close_1001_pos = source.find('await websocket.close(code=1001)') diff --git a/backend/tests/unit/test_redis_db_cache_serialization.py b/backend/tests/unit/test_redis_db_cache_serialization.py index 10374d62722..07a35c6dbe6 100644 --- a/backend/tests/unit/test_redis_db_cache_serialization.py +++ b/backend/tests/unit/test_redis_db_cache_serialization.py @@ -3,6 +3,7 @@ from __future__ import annotations from typing import Any, Dict, List, Optional +from unittest.mock import MagicMock import pytest @@ -13,8 +14,11 @@ class _FakeRedis: def __init__(self) -> None: self._store: Dict[str, Any] = {} - def set(self, key: str, value: Any, ex: Optional[int] = None) -> None: + def set(self, key: str, value: Any, ex: Optional[int] = None, nx: bool = False) -> Optional[bool]: + if nx and key in self._store: + return None self._store[key] = value + return True def get(self, key: str) -> Optional[Any]: return self._store.get(key) @@ -25,6 +29,15 @@ def expire(self, key: str, ttl: int) -> None: def mget(self, keys: List[str]) -> List[Optional[Any]]: return [self._store.get(key) for key in keys] + def eval(self, script: str, _numkeys: int, key: str, expected: str, *args: Any) -> int: + if self._store.get(key) != expected: + return 0 + if "redis.call('DEL'" in script: + del self._store[key] + else: + self._store[key] = 'done' + return 1 + @pytest.fixture def fake_redis(monkeypatch: pytest.MonkeyPatch) -> _FakeRedis: @@ -104,3 +117,73 @@ def test_apps_reviews_batch_round_trip(fake_redis: _FakeRedis) -> None: "app-b": {"uid-2": {"rating": 5}}, "app-missing": {}, } + + +def test_pusher_delivery_lease_reaches_done_with_a_bounded_key(fake_redis: _FakeRedis) -> None: + delivery_id = f"delivery-{'x' * 512}" + + state, lease_token = redis_db.begin_pusher_delivery("uid-1", delivery_id, redis_client=fake_redis) + assert state == 'claimed' + assert lease_token + assert redis_db.begin_pusher_delivery("uid-1", delivery_id, redis_client=fake_redis) == ('busy', None) + assert ( + redis_db.complete_pusher_delivery( + "uid-1", + delivery_id, + lease_token, + redis_client=fake_redis, + ) + is True + ) + assert redis_db.begin_pusher_delivery("uid-1", delivery_id, redis_client=fake_redis) == ('done', None) + assert redis_db.begin_pusher_delivery("uid-2", delivery_id, redis_client=fake_redis)[0] == 'claimed' + assert all(delivery_id not in key for key in fake_redis._store) + assert all(len(key) < 128 for key in fake_redis._store) + + +def test_pusher_failed_effect_releases_only_its_own_lease(fake_redis: _FakeRedis) -> None: + state, lease_token = redis_db.begin_pusher_delivery("uid-1", "delivery-1", redis_client=fake_redis) + assert state == 'claimed' + assert lease_token + assert ( + redis_db.abandon_pusher_delivery( + "uid-1", + "delivery-1", + "wrong-token", + redis_client=fake_redis, + ) + is False + ) + assert ( + redis_db.abandon_pusher_delivery( + "uid-1", + "delivery-1", + lease_token, + redis_client=fake_redis, + ) + is True + ) + assert redis_db.begin_pusher_delivery("uid-1", "delivery-1", redis_client=fake_redis)[0] == 'claimed' + + +def test_pusher_delivery_lease_fails_open_when_redis_is_unavailable() -> None: + class _UnavailableRedis: + def set(self, *args: Any, **kwargs: Any) -> None: + raise ConnectionError("redis unavailable") + + client = _UnavailableRedis() + + assert redis_db.begin_pusher_delivery("uid-1", "delivery-1", redis_client=client) == ('unavailable', None) + assert redis_db.complete_pusher_delivery("uid-1", "delivery-1", "lease-1", redis_client=client) is False + assert redis_db.abandon_pusher_delivery("uid-1", "delivery-1", "lease-1", redis_client=client) is False + + +def test_pusher_delivery_client_has_bounded_network_timeouts(monkeypatch: pytest.MonkeyPatch) -> None: + constructor = MagicMock(return_value=object()) + monkeypatch.setattr(redis_db.redis, 'Redis', constructor) + monkeypatch.setattr(redis_db, '_pusher_delivery_r', None) + + assert redis_db._get_pusher_delivery_redis() is constructor.return_value + assert constructor.call_args.kwargs['socket_connect_timeout'] == 0.5 + assert constructor.call_args.kwargs['socket_timeout'] == 0.5 + assert constructor.call_args.kwargs['retry_on_timeout'] is False diff --git a/backend/tests/unit/test_speaker_identification_delivery.py b/backend/tests/unit/test_speaker_identification_delivery.py new file mode 100644 index 00000000000..242ea89fdb9 --- /dev/null +++ b/backend/tests/unit/test_speaker_identification_delivery.py @@ -0,0 +1,238 @@ +import pytest +from unittest.mock import AsyncMock, MagicMock + +import utils.speaker_identification as speaker_identification +from utils.other import storage as storage_utils + + +async def _inline_run_blocking(_executor, func, *args, **kwargs): + return func(*args, **kwargs) + + +@pytest.fixture +def anyio_backend(): + return 'asyncio' + + +@pytest.mark.anyio +async def test_missing_audio_metadata_is_retryable_instead_of_false_success(monkeypatch): + monkeypatch.setattr(speaker_identification, 'run_blocking', _inline_run_blocking) + monkeypatch.setattr(speaker_identification.users_db, 'get_person', lambda _uid, _person_id: {}) + monkeypatch.setattr( + speaker_identification.users_db, + 'get_person_speech_samples_count', + lambda _uid, _person_id: 0, + ) + monkeypatch.setattr( + speaker_identification.conversations_db, + 'get_conversation', + lambda _uid, _conversation_id: { + 'started_at': 1.0, + 'transcript_segments': [{'id': 'segment-1', 'start': 0.0, 'end': 10.0}], + 'audio_files': [], + }, + ) + + result = await speaker_identification.extract_speaker_samples( + uid='uid-1', + person_id='person-1', + conversation_id='conversation-1', + segment_ids=['segment-1'], + ) + + assert result.status == 'retryable' + assert result.reason == 'audio_files_not_ready' + + +@pytest.mark.anyio +async def test_missing_requested_segment_is_retryable_instead_of_false_success(monkeypatch): + monkeypatch.setattr(speaker_identification, 'run_blocking', _inline_run_blocking) + monkeypatch.setattr(speaker_identification.users_db, 'get_person', lambda _uid, _person_id: {}) + monkeypatch.setattr( + speaker_identification.users_db, + 'get_person_speech_samples_count', + lambda _uid, _person_id: 0, + ) + monkeypatch.setattr( + speaker_identification.conversations_db, + 'get_conversation', + lambda _uid, _conversation_id: { + 'started_at': 1.0, + 'transcript_segments': [], + 'audio_files': [{'chunk_timestamps': [1.0]}], + }, + ) + + result = await speaker_identification.extract_speaker_samples( + uid='uid-1', + person_id='person-1', + conversation_id='conversation-1', + segment_ids=['segment-not-persisted-yet'], + ) + + assert result.status == 'retryable' + assert result.reason == 'transcript_segments_not_ready' + + +@pytest.mark.anyio +async def test_unhandled_extraction_error_is_retryable_instead_of_false_success(monkeypatch): + monkeypatch.setattr(speaker_identification, 'run_blocking', _inline_run_blocking) + + def fail(_uid, _person_id): + raise RuntimeError('database unavailable') + + monkeypatch.setattr(speaker_identification.users_db, 'get_person', fail) + + result = await speaker_identification.extract_speaker_samples( + uid='uid-1', + person_id='person-1', + conversation_id='conversation-1', + segment_ids=['segment-1'], + ) + + assert result.status == 'retryable' + assert result.reason == 'extraction_failed' + + +@pytest.mark.anyio +async def test_missing_person_is_retryable_before_upload_side_effects(monkeypatch): + monkeypatch.setattr(speaker_identification, 'run_blocking', _inline_run_blocking) + monkeypatch.setattr(speaker_identification.users_db, 'get_person', lambda _uid, _person_id: None) + get_conversation = MagicMock() + upload_sample = MagicMock() + monkeypatch.setattr(speaker_identification.conversations_db, 'get_conversation', get_conversation) + monkeypatch.setattr(speaker_identification, 'upload_person_speech_sample_from_bytes', upload_sample) + + result = await speaker_identification.extract_speaker_samples( + uid='uid-1', + person_id='person-1', + conversation_id='conversation-1', + segment_ids=['segment-1'], + ) + + assert result.status == 'retryable' + assert result.reason == 'person_not_ready' + get_conversation.assert_not_called() + upload_sample.assert_not_called() + + +@pytest.mark.anyio +async def test_sample_append_enforces_the_single_sample_limit_transactionally(monkeypatch): + monkeypatch.setattr(speaker_identification, 'run_blocking', _inline_run_blocking) + monkeypatch.setattr(speaker_identification.users_db, 'get_person', lambda _uid, _person_id: {}) + monkeypatch.setattr( + speaker_identification.users_db, + 'get_person_speech_samples_count', + lambda _uid, _person_id: 0, + ) + monkeypatch.setattr( + speaker_identification.conversations_db, + 'get_conversation', + lambda _uid, _conversation_id: { + 'started_at': 1_000.0, + 'transcript_segments': [ + { + 'id': 'segment-1', + 'start': 0.0, + 'end': 10.0, + 'text': 'hello there', + 'speaker_id': 1, + } + ], + 'audio_files': [{'chunk_timestamps': [1_000.0]}], + }, + ) + monkeypatch.setattr( + speaker_identification, + 'download_audio_chunks_and_merge', + lambda *_args, **_kwargs: b'audio', + ) + monkeypatch.setattr( + speaker_identification, + '_trim_pcm_audio', + lambda *_args, **_kwargs: b'\x00\x00' * (16_000 * 10), + ) + monkeypatch.setattr( + speaker_identification, + 'verify_and_transcribe_sample', + AsyncMock(return_value=('hello there', True, '')), + ) + upload_sample = MagicMock(return_value='users/uid-1/people/person-1/sample.pcm') + monkeypatch.setattr(speaker_identification, 'upload_person_speech_sample_from_bytes', upload_sample) + add_sample = MagicMock(return_value=True) + monkeypatch.setattr(speaker_identification.users_db, 'add_person_speech_sample', add_sample) + monkeypatch.setattr( + speaker_identification, + 'extract_embedding_from_bytes', + MagicMock(side_effect=RuntimeError('embedding unavailable')), + ) + + result = await speaker_identification.extract_speaker_samples( + uid='uid-1', + person_id='person-1', + conversation_id='conversation-1', + segment_ids=['segment-1'], + sample_rate=16_000, + delivery_id='delivery-1', + ) + + assert result.status == 'stored' + upload_sample.assert_called_once_with( + b'\x00\x00' * (16_000 * 10), + 'uid-1', + 'person-1', + 16_000, + 'speaker-sample\0uid-1\0person-1\0delivery-1', + ) + add_sample.assert_called_once_with( + 'uid-1', + 'person-1', + 'users/uid-1/people/person-1/sample.pcm', + transcript='hello there', + max_samples=1, + ) + + +def test_speech_sample_upload_reuses_one_hashed_object_for_a_stable_delivery(monkeypatch): + uploaded_paths = [] + + class FakeBlob: + def __init__(self, path): + self.path = path + + def upload_from_string(self, _payload, *, content_type): + assert content_type == 'audio/wav' + uploaded_paths.append(self.path) + + class FakeBucket: + def blob(self, path): + return FakeBlob(path) + + monkeypatch.setattr(storage_utils, '_get_speech_profiles_bucket', lambda **_kwargs: FakeBucket()) + + stable_key = 'speaker-sample\0uid-1\0person-1\0delivery-1' + first = storage_utils.upload_person_speech_sample_from_bytes( + b'\x00\x00', + 'uid-1', + 'person-1', + deduplication_key=stable_key, + ) + replay = storage_utils.upload_person_speech_sample_from_bytes( + b'\x00\x00', + 'uid-1', + 'person-1', + deduplication_key=stable_key, + ) + another = storage_utils.upload_person_speech_sample_from_bytes( + b'\x00\x00', + 'uid-1', + 'person-1', + deduplication_key='speaker-sample\0uid-1\0person-1\0delivery-2', + ) + legacy_first = storage_utils.upload_person_speech_sample_from_bytes(b'\x00\x00', 'uid-1', 'person-1') + legacy_second = storage_utils.upload_person_speech_sample_from_bytes(b'\x00\x00', 'uid-1', 'person-1') + + assert first == replay, 'one logical speaker delivery must overwrite one stable GCS object' + assert another != first + assert legacy_first != legacy_second + assert uploaded_paths == [first, replay, another, legacy_first, legacy_second] diff --git a/backend/tests/unit/test_task_integration_due_date_validation.py b/backend/tests/unit/test_task_integration_due_date_validation.py index a37ed216964..7482bce6786 100644 --- a/backend/tests/unit/test_task_integration_due_date_validation.py +++ b/backend/tests/unit/test_task_integration_due_date_validation.py @@ -72,3 +72,33 @@ def test_missing_due_date_is_allowed(app_client, monkeypatch): assert resp.status_code == 200 assert created.await_args.kwargs["due_date"] is None + + +def test_provider_failure_metadata_is_preserved_in_public_response(app_client, monkeypatch): + client, ti = app_client + monkeypatch.setattr(ti.users_db, "get_task_integration", _connected) + monkeypatch.setattr( + ti, + "create_task_internal", + AsyncMock( + return_value={ + "success": False, + "error": "ReadTimeout", + "error_code": "transport_error", + "retryable": False, + "ambiguous": True, + } + ), + ) + + resp = client.post("/v1/task-integrations/todoist/tasks", json={"title": "ambiguous task"}) + + assert resp.status_code == 200 + assert resp.json() == { + "success": False, + "external_task_id": None, + "error": "ReadTimeout", + "error_code": "transport_error", + "retryable": False, + "ambiguous": True, + } diff --git a/backend/tests/unit/test_task_integrations_ops.py b/backend/tests/unit/test_task_integrations_ops.py index 2039c396164..6300bb80a31 100644 --- a/backend/tests/unit/test_task_integrations_ops.py +++ b/backend/tests/unit/test_task_integrations_ops.py @@ -71,6 +71,7 @@ async def test_create_task_todoist_api_error_marks_disconnected(): assert result["success"] is False assert result["error_code"] == "api_error" + assert result["retryable"] is False mock_run_blocking.assert_awaited_once() saved = mock_run_blocking.call_args[0][4] assert saved["connected"] is False @@ -88,9 +89,145 @@ async def test_create_task_missing_access_token(): "success": False, "error": "No access token for todoist", "error_code": "no_access_token", + "retryable": False, + "ambiguous": False, } +@pytest.mark.asyncio +async def test_create_task_todoist_server_failure_is_ambiguous_not_retryable(): + client = AsyncMock(spec=httpx.AsyncClient) + client.post.return_value = _mock_response(503, text="Unavailable") + + result = await ops.create_task_internal( + uid="uid-4", + app_key="todoist", + integration={"connected": True, "access_token": "token"}, + title="Retry task", + client=client, + ) + + assert result["success"] is False + assert result["status_code"] == 503 + assert result["retryable"] is False + assert result["ambiguous"] is True + + +@pytest.mark.asyncio +async def test_create_task_todoist_rate_limit_is_safe_to_retry(): + client = AsyncMock(spec=httpx.AsyncClient) + client.post.return_value = _mock_response(429, text="Rate limited") + + result = await ops.create_task_internal( + uid="uid-4", + app_key="todoist", + integration={"connected": True, "access_token": "token"}, + title="Retry task", + client=client, + ) + + assert result["retryable"] is True + assert result["ambiguous"] is False + + +@pytest.mark.asyncio +async def test_create_task_todoist_missing_identity_is_ambiguous_not_string_none(): + client = AsyncMock(spec=httpx.AsyncClient) + client.post.return_value = _mock_response(201, {}) + + result = await ops.create_task_internal( + uid="uid-4", + app_key="todoist", + integration={"connected": True, "access_token": "token"}, + title="Ambiguous task", + client=client, + ) + + assert result == { + "success": False, + "error": "Provider response omitted task identity", + "error_code": "invalid_provider_response", + "retryable": False, + "ambiguous": True, + }, "provider success without task identity must be reported as ambiguous" + + +@pytest.mark.asyncio +async def test_create_task_unexpected_success_status_is_ambiguous(): + client = AsyncMock(spec=httpx.AsyncClient) + client.post.return_value = _mock_response(202, {"id": "possibly-created"}) + + result = await ops.create_task_internal( + uid="uid-4", + app_key="todoist", + integration={"connected": True, "access_token": "token"}, + title="Ambiguous task", + client=client, + ) + + assert result == { + "success": False, + "error": "Todoist response did not contain a completed task", + "error_code": "invalid_provider_response", + "status_code": 202, + "retryable": False, + "ambiguous": True, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "error", + [ + httpx.ReadTimeout("response timed out"), + httpx.ReadError("response failed"), + httpx.WriteTimeout("request write timed out"), + httpx.WriteError("request write failed"), + ], +) +async def test_create_task_transport_failure_is_ambiguous_not_blindly_retryable(error): + client = AsyncMock(spec=httpx.AsyncClient) + client.post.side_effect = error + + result = await ops.create_task_internal( + uid="uid-4", + app_key="todoist", + integration={"connected": True, "access_token": "token"}, + title="Ambiguous task", + client=client, + ) + + assert result["error_code"] == "transport_error" + assert result["retryable"] is False + assert result["ambiguous"] is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "error", + [ + httpx.PoolTimeout("connection pool exhausted"), + httpx.ConnectTimeout("connection timed out"), + httpx.ConnectError("connection failed"), + ], +) +async def test_create_task_pre_send_transport_failure_is_safe_to_retry(error): + client = AsyncMock(spec=httpx.AsyncClient) + client.post.side_effect = error + + result = await ops.create_task_internal( + uid="uid-4", + app_key="todoist", + integration={"connected": True, "access_token": "token"}, + title="Retry task", + client=client, + ) + + assert result["error_code"] == "transport_error" + assert result["retryable"] is True + assert result["ambiguous"] is False + + @pytest.mark.asyncio async def test_asana_retry_reuses_injected_client_for_refresh_and_retry(): client = AsyncMock(spec=httpx.AsyncClient) diff --git a/backend/tests/unit/test_tools_rest_memory_runtime_adapter.py b/backend/tests/unit/test_tools_rest_memory_runtime_adapter.py index 7b778fddc5b..bde9f587acc 100644 --- a/backend/tests/unit/test_tools_rest_memory_runtime_adapter.py +++ b/backend/tests/unit/test_tools_rest_memory_runtime_adapter.py @@ -72,11 +72,6 @@ def get_memories_by_ids(self, *args, **kwargs): raise AssertionError('legacy get_memories_by_ids must not run for memory denied/enabled tools REST reads') -class _UnexpectedLegacyVectorDb: - def find_similar_memories(self, *args, **kwargs): - raise AssertionError('legacy vector search must not run for memory denied/enabled tools REST reads') - - def test_tools_rest_get_memories_text_requests_legacy_safe_memory_decision(monkeypatch): memory_services = _load_memory_services() captured = [] @@ -168,13 +163,6 @@ def fake_search_adapter(**kwargs): ) monkeypatch.setattr(memory_services, 'memory_db', _UnexpectedLegacyMemoryDb()) - monkeypatch.setattr(memory_services, 'vector_db', _UnexpectedLegacyVectorDb()) - monkeypatch.setattr(memory_services.notification_db, 'get_user_time_zone', lambda _uid: 'UTC') - monkeypatch.setattr( - memory_services, - 'pin_memory_system', - lambda *_args, **_kwargs: memory_services.MemorySystem.LEGACY, - ) monkeypatch.setattr( memory_services, 'search_memory_default_chat_memories_vector_decision_text', fake_search_adapter ) @@ -200,13 +188,6 @@ def fake_search_adapter(**kwargs): def test_tools_rest_search_memories_text_preserves_adapter_denied_or_empty_memory_states(monkeypatch): memory_services = _load_memory_services() monkeypatch.setattr(memory_services, 'memory_db', _UnexpectedLegacyMemoryDb()) - monkeypatch.setattr(memory_services, 'vector_db', _UnexpectedLegacyVectorDb()) - monkeypatch.setattr(memory_services.notification_db, 'get_user_time_zone', lambda _uid: 'UTC') - monkeypatch.setattr( - memory_services, - 'pin_memory_system', - lambda *_args, **_kwargs: memory_services.MemorySystem.LEGACY, - ) monkeypatch.setattr( memory_services, diff --git a/backend/tests/unit/test_tools_router.py b/backend/tests/unit/test_tools_router.py index 7f4c9238a12..154b5b6e200 100644 --- a/backend/tests/unit/test_tools_router.py +++ b/backend/tests/unit/test_tools_router.py @@ -4,7 +4,7 @@ 1. get_conversations_text — date parsing, limit caps, empty results 2. search_conversations_text — query routing, date conversion to timestamps 3. get_memories_text — date parsing, locked memory filtering -4. search_memories_text — vector search delegation +4. search_memories_text — shared legacy search delegation 5. get_action_items_text — date parsing, status filtering 6. create_action_item_text — validation, default due date, past-date rejection 7. update_action_item_text — exists check, field updates @@ -177,7 +177,6 @@ def _load_module_from_file(module_name, file_path): # Stub database.vector_db vector_db = _stub_module("database.vector_db") vector_db.query_vectors = MagicMock(return_value=[]) -vector_db.find_similar_memories = MagicMock(return_value=[]) # Stub database.action_items action_items_db = _stub_module("database.action_items") @@ -229,6 +228,7 @@ def _load_module_from_file(module_name, file_path): memory_service_stub = _stub_module("utils.memory.memory_service") memory_service_stub.MemoryService = MagicMock +memory_service_stub.search_legacy_memories = MagicMock(return_value=[]) surface_routing_stub = _stub_module("utils.memory.surface_routing") surface_routing_stub.pin_memory_system = MagicMock(return_value=memory_system_stub.MemorySystem.LEGACY) @@ -650,10 +650,9 @@ def test_limit_cap(self): # =========================================================================== class TestSearchMemoriesText: def setup_method(self): - vector_db.find_similar_memories.reset_mock() - vector_db.find_similar_memories.return_value = [] - memories_db.get_memories_by_ids.reset_mock() - memories_db.get_memories_by_ids.return_value = [] + memories_svc.search_legacy_memories.reset_mock() + memories_svc.search_legacy_memories.return_value = [] + memories_svc.search_legacy_memories.side_effect = None def test_no_results(self): result = memories_svc.search_memories_text(uid="test-uid", query="cooking") @@ -661,15 +660,16 @@ def test_no_results(self): assert "cooking" in result def test_with_results(self): - vector_db.find_similar_memories.return_value = [ - {'memory_id': 'mem-1', 'score': 0.95}, - ] - memories_db.get_memories_by_ids.return_value = [ - {'id': 'mem-1', 'content': 'likes pasta', 'is_locked': False, 'created_at': datetime.now(timezone.utc)}, + memories_svc.search_legacy_memories.return_value = [ + types.SimpleNamespace( + memory=FakeMemoryDB(id='mem-1', content='likes pasta', created_at=datetime.now(timezone.utc)), + score=0.95, + ) ] result = memories_svc.search_memories_text(uid="test-uid", query="food") assert "likes pasta" in result assert "0.95" in result + memories_svc.search_legacy_memories.assert_called_once_with("test-uid", "food", limit=5) # =========================================================================== @@ -852,8 +852,9 @@ def setup_app(self): memories_db.get_memories.return_value = [] memories_db.get_memories_by_ids.reset_mock() memories_db.get_memories_by_ids.return_value = [] - vector_db.find_similar_memories.reset_mock() - vector_db.find_similar_memories.return_value = [] + memories_svc.search_legacy_memories.reset_mock() + memories_svc.search_legacy_memories.return_value = [] + memories_svc.search_legacy_memories.side_effect = None action_items_db.get_action_items.reset_mock() action_items_db.get_action_items.return_value = [] action_items_db.get_action_item.reset_mock() @@ -1090,36 +1091,28 @@ def test_search_conversations_all_locked_returns_empty(self): # =========================================================================== -# Tests: Search memories locked filtering +# Tests: Search memories shared helper results # =========================================================================== -class TestSearchMemoriesLockedFiltering: +class TestSearchMemoriesSharedHelperResults: def setup_method(self): - vector_db.find_similar_memories.reset_mock() - memories_db.get_memories_by_ids.reset_mock() - - def test_search_memories_filters_locked(self): - """Locked memories are excluded from search results.""" - vector_db.find_similar_memories.return_value = [ - {'memory_id': 'mem-1', 'score': 0.9}, - {'memory_id': 'mem-2', 'score': 0.8}, - ] - memories_db.get_memories_by_ids.return_value = [ - {'id': 'mem-1', 'content': 'visible', 'is_locked': False, 'created_at': datetime.now(timezone.utc)}, - {'id': 'mem-2', 'content': 'locked', 'is_locked': True, 'created_at': datetime.now(timezone.utc)}, + memories_svc.search_legacy_memories.reset_mock() + memories_svc.search_legacy_memories.return_value = [] + memories_svc.search_legacy_memories.side_effect = None + + def test_search_memories_formats_shared_helper_results(self): + """The tool formats the visible matches returned by the shared helper.""" + memories_svc.search_legacy_memories.return_value = [ + types.SimpleNamespace( + memory=FakeMemoryDB(id='mem-1', content='visible', created_at=datetime.now(timezone.utc)), + score=0.9, + ) ] result = memories_svc.search_memories_text(uid="test-uid", query="test") assert "visible" in result - assert "locked" not in result assert "1 memories" in result - def test_search_memories_all_locked_returns_empty(self): - """All locked memories in search returns 'no memories' message.""" - vector_db.find_similar_memories.return_value = [ - {'memory_id': 'mem-1', 'score': 0.9}, - ] - memories_db.get_memories_by_ids.return_value = [ - {'id': 'mem-1', 'content': 'locked', 'is_locked': True, 'created_at': datetime.now(timezone.utc)}, - ] + def test_search_memories_empty_shared_helper_result_returns_empty(self): + """Filtering everything in the shared helper returns the no-results message.""" result = memories_svc.search_memories_text(uid="test-uid", query="test") assert "No memories found" in result @@ -1133,7 +1126,9 @@ def setup_method(self): conversations_db.get_conversations_by_id.reset_mock() vector_db.query_vectors.reset_mock() memories_db.get_memories.reset_mock() - vector_db.find_similar_memories.reset_mock() + memories_svc.search_legacy_memories.reset_mock() + memories_svc.search_legacy_memories.return_value = [] + memories_svc.search_legacy_memories.side_effect = None action_items_db.get_action_items.reset_mock() action_items_db.create_action_item.reset_mock() action_items_db.get_action_item.reset_mock() @@ -1163,13 +1158,13 @@ def test_get_memories_db_error(self): assert "Firestore down" in result memories_db.get_memories.side_effect = None - def test_search_memories_vector_error(self): - """Vector DB failure in search_memories returns error text.""" - vector_db.find_similar_memories.side_effect = Exception("Vector timeout") + def test_search_memories_shared_helper_error(self): + """A shared legacy search failure returns error text.""" + memories_svc.search_legacy_memories.side_effect = Exception("Vector timeout") result = memories_svc.search_memories_text(uid="test-uid", query="test") assert "Error" in result assert "Vector timeout" in result - vector_db.find_similar_memories.side_effect = None + memories_svc.search_legacy_memories.side_effect = None def test_get_action_items_db_error(self): """DB failure in get_action_items returns error text.""" @@ -1222,8 +1217,9 @@ def setup_method(self): vector_db.query_vectors.return_value = [] memories_db.get_memories.reset_mock() memories_db.get_memories.return_value = [] - vector_db.find_similar_memories.reset_mock() - vector_db.find_similar_memories.return_value = [] + memories_svc.search_legacy_memories.reset_mock() + memories_svc.search_legacy_memories.return_value = [] + memories_svc.search_legacy_memories.side_effect = None action_items_db.get_action_items.reset_mock() action_items_db.get_action_items.return_value = [] @@ -1236,9 +1232,7 @@ def test_search_conversations_limit_cap(self): def test_search_memories_limit_cap(self): """search_memories_text caps limit at 20.""" memories_svc.search_memories_text(uid="test-uid", query="test", limit=100) - call_kwargs = vector_db.find_similar_memories.call_args - # limit is positional arg 3 or keyword - assert call_kwargs[1].get('limit', call_kwargs[0][2] if len(call_kwargs[0]) > 2 else 20) <= 20 + memories_svc.search_legacy_memories.assert_called_once_with('test-uid', 'test', limit=20) def test_action_items_limit_cap(self): """get_action_items_text caps limit at 500.""" diff --git a/backend/tests/unit/test_users_webhook_url_validation.py b/backend/tests/unit/test_users_webhook_url_validation.py index 5a3ca447510..f0ab204b8b5 100644 --- a/backend/tests/unit/test_users_webhook_url_validation.py +++ b/backend/tests/unit/test_users_webhook_url_validation.py @@ -121,13 +121,42 @@ def test_missing_url_returns_422(): SetUserWebhookUrlRequest() -def test_valid_url_sets(): - with patch.object(users_mod, 'set_user_webhook_db') as setdb, patch.object(users_mod, 'disable_user_webhook_db'): +def test_nonempty_url_reenables_delivery_and_resets_failure_health(): + wtype = users_mod.WebhookType.audio_bytes + webhook_url = 'https://example.com/audio?token=replacement' + with ( + patch.object(users_mod, 'set_user_webhook_db') as setdb, + patch.object(users_mod, 'get_user_webhook_db', return_value=webhook_url), + patch.object(users_mod, 'disable_user_webhook_db') as disable, + patch.object(users_mod, 'enable_user_webhook_db') as enable, + patch.object(users_mod, 'reset_user_webhook_delivery_health') as reset_health, + ): result = users_mod.set_user_webhook_endpoint( - wtype='audio_bytes', data=SetUserWebhookUrlRequest(url='http://x'), uid='u1' + wtype=wtype, data=SetUserWebhookUrlRequest(url=webhook_url), uid='u1' ) + + assert result['status'] == 'ok' + setdb.assert_called_once_with('u1', wtype, webhook_url) + disable.assert_not_called() + enable.assert_called_once_with('u1', wtype) + reset_health.assert_called_once_with('u1', wtype, webhook_url) + + +def test_empty_url_disables_without_resetting_failure_health(): + wtype = users_mod.WebhookType.audio_bytes + with ( + patch.object(users_mod, 'set_user_webhook_db') as setdb, + patch.object(users_mod, 'disable_user_webhook_db') as disable, + patch.object(users_mod, 'enable_user_webhook_db') as enable, + patch.object(users_mod, 'reset_user_webhook_delivery_health') as reset_health, + ): + result = users_mod.set_user_webhook_endpoint(wtype=wtype, data=SetUserWebhookUrlRequest(url=''), uid='u1') + assert result['status'] == 'ok' - setdb.assert_called_once() + setdb.assert_called_once_with('u1', wtype, '') + disable.assert_called_once_with('u1', wtype) + reset_health.assert_not_called() + enable.assert_not_called() def test_get_missing_webhook_url_validates_as_nullable_response(): diff --git a/backend/tests/unit/test_webhook_auto_disable.py b/backend/tests/unit/test_webhook_auto_disable.py index 4a04b0fca8c..7ce642650a8 100644 --- a/backend/tests/unit/test_webhook_auto_disable.py +++ b/backend/tests/unit/test_webhook_auto_disable.py @@ -15,6 +15,7 @@ import httpx import pytest +import database.webhook_health as webhook_health_db from testing.import_isolation import load_module_fresh, stub_modules from utils.apps import validate_app_endpoints_for_reenable @@ -318,7 +319,7 @@ async def test_dev_webhook_disabled_on_threshold(self): patch("utils.webhooks.record_dev_webhook_failure", return_value=True) as mock_fail, patch("utils.webhooks.disable_user_webhook_db") as mock_disable, patch("utils.webhooks.send_notification") as mock_notify, - patch("utils.webhooks._DEV_WEBHOOK_RETRY_DELAYS", ()), + patch("utils.webhooks._REALTIME_DEV_WEBHOOK_RETRY_DELAYS", ()), ): await realtime_transcript_webhook("uid-1", [{"text": "hello"}]) mock_fail.assert_called_once() @@ -367,7 +368,7 @@ async def test_dev_webhook_exception_records_failure(self): patch("utils.webhooks.get_webhook_client", return_value=mock_client), patch("utils.webhooks.get_webhook_circuit_breaker", return_value=mock_cb), patch("utils.webhooks.record_dev_webhook_failure", return_value=False) as mock_fail, - patch("utils.webhooks._DEV_WEBHOOK_RETRY_DELAYS", ()), + patch("utils.webhooks._REALTIME_DEV_WEBHOOK_RETRY_DELAYS", ()), ): await realtime_transcript_webhook("uid-1", [{"text": "hello"}]) mock_fail.assert_called_once() @@ -406,7 +407,7 @@ async def fake_sleep(delay): patch("utils.webhooks.get_webhook_semaphore", return_value=mock_sem), patch("utils.webhooks.record_dev_webhook_success") as mock_success, patch("utils.webhooks.record_dev_webhook_failure") as mock_fail, - patch("utils.webhooks._DEV_WEBHOOK_RETRY_DELAYS", (0.01,)), + patch("utils.webhooks._REALTIME_DEV_WEBHOOK_RETRY_DELAYS", (0.01,)), patch("utils.webhooks.asyncio.sleep", side_effect=fake_sleep), ): await realtime_transcript_webhook("uid-1", [{"text": "hello"}]) @@ -1379,18 +1380,17 @@ def test_non_disabled_app_loaded(self): class TestDevWebhookManualReEnable: """Test that manual dev webhook re-enable clears health state.""" - def test_success_on_enable_clears_state(self): - """record_dev_webhook_success called on manual enable should reset all fields.""" - from database.webhook_health import record_dev_webhook_success - + def test_reset_on_enable_clears_state_without_faking_success(self): mock_r = MagicMock() - with patch("database.webhook_health.r", mock_r): - record_dev_webhook_success("uid-1", "realtime_transcript") + with patch.object(webhook_health_db, 'r', mock_r): + webhook_health_db.reset_dev_webhook_health("uid-1", "realtime_transcript") mapping = mock_r.hset.call_args.kwargs.get('mapping') or mock_r.hset.call_args[1].get('mapping') assert mapping['failure_count'] == '0' assert mapping['disabled'] == '0' assert mapping['last_error'] == '' + assert mapping['last_success_at'] == '' + assert mapping['last_status'] == '' mock_r.expire.assert_called_once() diff --git a/backend/tests/unit/test_x_memory_extraction_retry.py b/backend/tests/unit/test_x_memory_extraction_retry.py index e0b1bf4a220..a41d6731972 100644 --- a/backend/tests/unit/test_x_memory_extraction_retry.py +++ b/backend/tests/unit/test_x_memory_extraction_retry.py @@ -96,5 +96,26 @@ def write(self, *_args, **_kwargs): assert acknowledgements == [('uid-1', ['post-1'])] +def test_pending_x_source_remains_unacknowledged_after_partial_legacy_projection(monkeypatch): + post = {'id': 'post-1', 'text': 'I prefer tea', 'created_at': '2026-07-14T00:00:00Z', 'kind': 'tweet'} + acknowledgements = [] + memory = Memory(content='User prefers tea', category=MemoryCategory.interesting) + + monkeypatch.setattr(x_connector, 'extract_memories_from_text', lambda *args: [memory]) + monkeypatch.setattr(x_connector, 'resolve_memory_system', lambda *args, **kwargs: MemorySystem.LEGACY) + monkeypatch.setattr(x_connector.memories_db, 'save_memories', lambda *args, **kwargs: None) + monkeypatch.setattr(x_connector, 'upsert_memory_vectors_batch', lambda *args, **kwargs: 0) + monkeypatch.setattr( + x_connector.x_posts_db, + 'mark_memory_extraction_completed', + lambda uid, post_ids: acknowledgements.append((uid, post_ids)), + ) + + with pytest.raises(RuntimeError, match='partial expected=1 actual=0'): + x_connector._extract_and_index('uid-1', [post]) + + assert acknowledgements == [] + + async def _async_value(value): return value diff --git a/backend/tests/unit/utils/test_listen_pusher_session.py b/backend/tests/unit/utils/test_listen_pusher_session.py index 14974234b2a..9196e887c26 100644 --- a/backend/tests/unit/utils/test_listen_pusher_session.py +++ b/backend/tests/unit/utils/test_listen_pusher_session.py @@ -4,6 +4,7 @@ import pytest +import utils.listen_pusher_session as listen_pusher_module from utils.listen_pusher_session import ( TARGET_SAMPLE_RATE, ListenPusherSession, @@ -13,11 +14,12 @@ class FakePusherWebSocket: - def __init__(self, incoming=None): + def __init__(self, incoming=None, *, ack_supported=False): self.sent = [] self.incoming = list(incoming or []) self.closed_codes = [] self.on_recv = None + self.response_headers = {'X-Omi-Delivery-Ack': '1'} if ack_supported else {} async def send(self, data): self.sent.append(bytes(data)) @@ -33,6 +35,66 @@ async def close(self, code=1000): self.closed_codes.append(code) +class FailFirstFrameWebSocket(FakePusherWebSocket): + def __init__(self, failing_frame_type: int, *, block_failure: bool = False): + super().__init__() + self.failing_frame_type = failing_frame_type + self.block_failure = block_failure + self.failure_started = asyncio.Event() + self.release_failure = asyncio.Event() + self.failed = False + + async def send(self, data): + frame = bytes(data) + if frame_type(frame) == self.failing_frame_type and not self.failed: + self.failed = True + self.failure_started.set() + if self.block_failure: + await self.release_failure.wait() + raise RuntimeError(f"failed frame {self.failing_frame_type}") + self.sent.append(frame) + + +class DeliverThenFailFirstFrameWebSocket(FakePusherWebSocket): + def __init__(self, failing_frame_type: int): + super().__init__() + self.failing_frame_type = failing_frame_type + self.failed = False + + async def send(self, data): + frame = bytes(data) + self.sent.append(frame) + if frame_type(frame) == self.failing_frame_type and not self.failed: + self.failed = True + raise RuntimeError(f"ambiguous frame {self.failing_frame_type}") + + +class BlockingConnectionClosedWebSocket(FakePusherWebSocket): + def __init__(self, operation: str): + super().__init__() + self.operation = operation + self.failure_started = asyncio.Event() + self.release_failure = asyncio.Event() + + async def recv(self): + if self.operation == "recv": + self.failure_started.set() + await self.release_failure.wait() + raise ControlledConnectionClosed() + return await super().recv() + + async def send(self, data): + if self.operation == "send": + self.failure_started.set() + await self.release_failure.wait() + raise ControlledConnectionClosed() + await super().send(data) + + +class ControlledConnectionClosed(Exception): + pass + + def frame_type(frame: bytes) -> int: return struct.unpack("I", frame[:4])[0] @@ -58,6 +120,11 @@ def fenced_response_201(conversation_id: str): return struct.pack(" bool: return status_code >= 500 or status_code in _RETRYABLE_DELIVERY_STATUSES +_REALTIME_APP_WEBHOOK_RETRY_DELAYS = (0.5, 2.0) + + +async def _post_realtime_app_webhook( + app_id: str, + webhook_url: str, + *, + idempotency_key: str | None = None, + retry_delays: tuple[float, ...] = _REALTIME_APP_WEBHOOK_RETRY_DELAYS, + **request_kwargs, +): + """Retry a realtime app delivery without changing its receiver-visible identity.""" + headers = dict(request_kwargs.pop('headers', {}) or {}) + headers.setdefault('X-Omi-Idempotency-Key', idempotency_key or str(uuid.uuid4())) + request_kwargs['headers'] = headers + client = get_webhook_client() + attempts = len(retry_delays) + 1 + last_response = None + last_exception: httpx.TransportError | None = None + + for attempt_index in range(attempts): + try: + async with get_webhook_semaphore(): + response = await client.post(webhook_url, **request_kwargs) + last_response = response + last_exception = None + if 200 <= response.status_code < 300: + return response + if not _delivery_failure_is_retryable(response.status_code): + return response + except httpx.TransportError as error: + last_response = None + last_exception = error + + if attempt_index < len(retry_delays): + delay = retry_delays[attempt_index] + logger.warning( + 'Realtime app webhook retry app=%s attempt=%s/%s delay=%ss', + app_id, + attempt_index + 1, + attempts, + f'{delay:g}', + ) + await asyncio.sleep(delay) + + if last_response is not None: + return last_response + if last_exception is not None: + raise last_exception + raise RuntimeError('Realtime app webhook failed without a response') + + def _notify_app_owner(app_id: str, title: str, body: str): """Send a push notification to the app owner about webhook health.""" try: @@ -342,10 +395,18 @@ async def trigger_realtime_integrations( segments: list[dict], conversation_id: str | None, source: str | None = None, + *, + idempotency_key: str | None = None, ): logger.info(f"trigger_realtime_integrations {uid}") """REALTIME STREAMING""" - return await _async_trigger_realtime_integrations(uid, segments, conversation_id, source=source) + return await _async_trigger_realtime_integrations( + uid, + segments, + conversation_id, + source=source, + idempotency_key=idempotency_key, + ) async def trigger_realtime_audio_bytes(uid: str, sample_rate: int, data: bytearray): @@ -779,6 +840,8 @@ async def _async_trigger_realtime_integrations( segments: List[dict], conversation_id: str | None, source: str | None = None, + *, + idempotency_key: str | None = None, ) -> dict: # Paywall: skip mentor + third-party proactive notifications when this # transcription session belongs to a paywalled desktop user. @@ -842,15 +905,15 @@ async def _single(app: App): return try: - async with get_webhook_semaphore(): - client = get_webhook_client() - response = await client.post( - pinned_url, - json={"session_id": uid, "segments": segments}, - headers=pin_kwargs['headers'], - extensions=pin_kwargs['extensions'], - follow_redirects=False, - ) + response = await _post_realtime_app_webhook( + app.id, + pinned_url, + json={"session_id": uid, "segments": segments}, + headers=pin_kwargs['headers'], + extensions=pin_kwargs['extensions'], + follow_redirects=False, + idempotency_key=idempotency_key, + ) if response.status_code < 200 or response.status_code >= 300: cb.record_failure() error_str = f'HTTP {response.status_code}' diff --git a/backend/utils/http_client.py b/backend/utils/http_client.py index b39b21d4541..de1f6536978 100644 --- a/backend/utils/http_client.py +++ b/backend/utils/http_client.py @@ -195,6 +195,13 @@ def record_failure(self): _CIRCUIT_BREAKER_IDLE_TTL = 3600 # seconds — evict entries idle for 1 hour +def _webhook_circuit_breaker_key(url: str) -> str: + try: + return url.split('?')[0].split('#')[0] + except (IndexError, AttributeError): + return url + + def get_webhook_circuit_breaker(url: str) -> WebhookCircuitBreaker: """Get or create a circuit breaker for a webhook target URL. @@ -202,11 +209,7 @@ def get_webhook_circuit_breaker(url: str) -> WebhookCircuitBreaker: different webhook endpoints on the same host are isolated from each other. Evicts stale entries when the registry grows beyond _CIRCUIT_BREAKER_MAX_ENTRIES. """ - try: - # Strip query params but keep scheme + host + path - key = url.split('?')[0].split('#')[0] - except (IndexError, AttributeError): - key = url + key = _webhook_circuit_breaker_key(url) if key not in _webhook_circuit_breakers: if len(_webhook_circuit_breakers) > _CIRCUIT_BREAKER_MAX_ENTRIES: _evict_stale_circuit_breakers() @@ -214,6 +217,11 @@ def get_webhook_circuit_breaker(url: str) -> WebhookCircuitBreaker: return _webhook_circuit_breakers[key] +def reset_webhook_circuit_breaker(url: str) -> None: + """Forget prior failures when a user explicitly replaces or re-enables a target.""" + _webhook_circuit_breakers.pop(_webhook_circuit_breaker_key(url), None) + + def _evict_stale_circuit_breakers(): """Remove circuit breaker entries not accessed for longer than _CIRCUIT_BREAKER_IDLE_TTL. diff --git a/backend/utils/listen_pusher_session.py b/backend/utils/listen_pusher_session.py index 775121264f4..60785de9e35 100644 --- a/backend/utils/listen_pusher_session.py +++ b/backend/utils/listen_pusher_session.py @@ -4,15 +4,17 @@ import random import struct import time +import uuid from collections import deque from dataclasses import dataclass from enum import Enum from typing import Any, Awaitable, Callable, cast, Deque, Dict, List, Optional, Tuple -from websockets.client import WebSocketClientProtocol +from websockets.legacy.client import WebSocketClientProtocol from websockets.exceptions import ConnectionClosed from utils.metrics import PUSHER_CIRCUIT_BREAKER_REJECTIONS, PUSHER_SESSION_DEGRADED +from utils.observability.fallback import record_fallback from utils.pusher import PusherCircuitBreakerOpen, connect_to_trigger_pusher # Typed wrapper because utils.pusher.connect_to_trigger_pusher uses the untyped @@ -41,6 +43,10 @@ class PusherReconnectState(str, Enum): PENDING_REQUEST_TIMEOUT = 120 MAX_RETRIES_PER_REQUEST = 3 PENDING_REQUEST_RECOVERY_COOLDOWN = 300 +PUSHER_DELIVERY_ACK_TIMEOUT = 5.0 +PUSHER_CLOSE_ACK_TIMEOUT = 10.0 +PUSHER_SOCKET_CLOSE_TIMEOUT = 2.0 +PUSHER_DELIVERY_DRAIN_OPCODE = 107 @dataclass @@ -72,27 +78,57 @@ class ListenPusherSessionDeps: monotonic: Callable[[], float] = time.monotonic +@dataclass +class _TranscriptDelivery: + delivery_id: str + conversation_id: Optional[str] + segments: List[Dict[str, Any]] + last_sent_ws: Optional[WebSocketClientProtocol] = None + last_sent_at: float = 0.0 + + +@dataclass +class _AudioDelivery: + delivery_id: str + conversation_id: Optional[str] + chunks: List[bytes] + last_received: Optional[float] + + class ListenPusherSession: def __init__(self, config: ListenPusherSessionConfig, deps: ListenPusherSessionDeps): self.config = config self.deps = deps self.pusher_ws: Optional[WebSocketClientProtocol] = None self.pusher_connect_lock = asyncio.Lock() + self.pusher_receive_lock = asyncio.Lock() self.pusher_connected = False + self.delivery_ack_supported = False self.reconnect_state = PusherReconnectState.CONNECTED self.reconnect_attempts = 0 self.reconnect_task: Optional[asyncio.Task[None]] = None self.degraded_since: float = 0.0 self.segment_buffers: Deque[Dict[str, Any]] = deque(maxlen=config.max_segment_buffer_size) + self.segment_buffer_conversation_ids: Deque[Optional[str]] = deque(maxlen=config.max_segment_buffer_size) + self.pending_transcript_delivery: Optional[_TranscriptDelivery] = None + self.transcript_flush_lock = asyncio.Lock() self.last_synced_conversation_id: Optional[str] = None self.pending_conversation_requests: Dict[str, Dict[str, Any]] = {} - self.pending_request_event = asyncio.Event() self.pending_speaker_sample_requests: Deque[Tuple[str, str, List[str]]] = deque( maxlen=config.max_pending_speaker_sample_requests ) + self.pending_speaker_sample_delivery_ids: Dict[Tuple[str, str, Tuple[str, ...]], str] = {} + self.pending_speaker_sample_sent: Dict[ + Tuple[str, str, Tuple[str, ...]], Tuple[WebSocketClientProtocol, float] + ] = {} + self.speaker_sample_send_lock = asyncio.Lock() self.audio_chunks: Deque[bytes] = deque() + self.audio_chunk_conversation_ids: Deque[Optional[str]] = deque() + self.audio_chunk_received_at: Deque[float] = deque() self.audio_total_size = 0 self.audio_buffer_last_received: Optional[float] = None + self.pending_audio_delivery: Optional[_AudioDelivery] = None + self.audio_flush_lock = asyncio.Lock() @property def uid(self): @@ -103,7 +139,150 @@ def session_id(self): return self.config.session_id def transcript_send(self, segments: List[Dict[str, Any]]) -> None: - self.segment_buffers.extend(segments) + conversation_id = self.deps.get_current_conversation_id() + pending_size = len(self.pending_transcript_delivery.segments) if self.pending_transcript_delivery else 0 + live_capacity = max(0, self.config.max_segment_buffer_size - pending_size) + dropped = False + for segment in segments: + if len(self.segment_buffers) >= live_capacity: + dropped = True + continue + self.segment_buffers.append(segment) + self.segment_buffer_conversation_ids.append(conversation_id) + if dropped: + record_fallback( + component='pusher', + from_mode='listen_transcript_buffer', + to_mode='drop_newest', + reason='capacity_full', + outcome='degraded', + log=logger, + ) + + def _delivery_due( + self, + last_sent_ws: Optional[WebSocketClientProtocol], + last_sent_at: float, + pusher_ws: WebSocketClientProtocol, + ) -> bool: + return last_sent_ws is not pusher_ws or self.deps.monotonic() - last_sent_at >= PUSHER_DELIVERY_ACK_TIMEOUT + + def _ack_pusher_delivery(self, delivery_id: str) -> None: + if self.pending_transcript_delivery is not None and self.pending_transcript_delivery.delivery_id == delivery_id: + self.pending_transcript_delivery = None + return + for request in list(self.pending_speaker_sample_requests): + person_id, conv_id, segment_ids = request + request_key = (person_id, conv_id, tuple(segment_ids)) + if self.pending_speaker_sample_delivery_ids.get(request_key) != delivery_id: + continue + self.pending_speaker_sample_requests.remove(request) + self.pending_speaker_sample_delivery_ids.pop(request_key, None) + self.pending_speaker_sample_sent.pop(request_key, None) + return + + async def _retry_pending_speaker_sample_requests(self) -> None: + if not self.pusher_connected or not self.pusher_ws or not self.pending_speaker_sample_requests: + return + async with self.speaker_sample_send_lock: + for person_id, conv_id, segment_ids in list(self.pending_speaker_sample_requests): + if not self.pusher_connected: + break + await self._send_speaker_sample_request(person_id, conv_id, segment_ids) + + def _take_transcript_delivery(self) -> Optional[_TranscriptDelivery]: + if self.pending_transcript_delivery is not None: + return self.pending_transcript_delivery + if not self.segment_buffers: + return None + + conversation_id = self.segment_buffer_conversation_ids[0] + segments: List[Dict[str, Any]] = [] + while self.segment_buffers and self.segment_buffer_conversation_ids[0] == conversation_id: + segments.append(self.segment_buffers.popleft()) + self.segment_buffer_conversation_ids.popleft() + delivery = _TranscriptDelivery( + delivery_id=str(uuid.uuid4()), + conversation_id=conversation_id or self.deps.get_current_conversation_id(), + segments=segments, + ) + self.pending_transcript_delivery = delivery + return delivery + + def _take_audio_delivery(self) -> Optional[_AudioDelivery]: + if self.pending_audio_delivery is not None: + return self.pending_audio_delivery + if not self.audio_chunks: + return None + + conversation_id = self.audio_chunk_conversation_ids[0] + chunks: List[bytes] = [] + received_at: List[float] = [] + total_size = 0 + while self.audio_chunks and self.audio_chunk_conversation_ids[0] == conversation_id: + chunk = self.audio_chunks.popleft() + self.audio_chunk_conversation_ids.popleft() + received_at.append(self.audio_chunk_received_at.popleft()) + chunks.append(chunk) + total_size += len(chunk) + + delivery = _AudioDelivery( + delivery_id=str(uuid.uuid4()), + conversation_id=conversation_id or self.deps.get_current_conversation_id(), + chunks=chunks, + last_received=received_at[-1] if received_at else None, + ) + self.audio_total_size -= total_size + self.audio_buffer_last_received = self.audio_chunk_received_at[-1] if self.audio_chunk_received_at else None + self.pending_audio_delivery = delivery + return delivery + + def _buffer_pending_speaker_sample_request( + self, + person_id: str, + conv_id: str, + segment_ids: List[str], + ) -> bool: + request = (person_id, conv_id, list(segment_ids)) + request_key = (person_id, conv_id, tuple(segment_ids)) + if request in self.pending_speaker_sample_requests: + return True + if self.config.max_pending_speaker_sample_requests <= 0: + record_fallback( + component='pusher', + from_mode='listen_speaker_sample_buffer', + to_mode='drop_newest', + reason='capacity_full', + outcome='degraded', + log=logger, + ) + return False + if len(self.pending_speaker_sample_requests) >= self.config.max_pending_speaker_sample_requests: + record_fallback( + component='pusher', + from_mode='listen_speaker_sample_buffer', + to_mode='drop_newest', + reason='capacity_full', + outcome='degraded', + log=logger, + ) + return False + self.pending_speaker_sample_requests.append(request) + self.pending_speaker_sample_delivery_ids[request_key] = str(uuid.uuid4()) + return True + + def _mark_failed_socket_disconnected( + self, + failed_ws: WebSocketClientProtocol, + *, + auto_reconnect: bool, + ) -> None: + if self.pusher_ws is not failed_ws: + return + if auto_reconnect: + self._mark_disconnected() + else: + self.pusher_connected = False def _buffer_pending_conversation_request( self, @@ -129,7 +308,6 @@ def _buffer_pending_conversation_request( 'finalization_job_id': finalization_job_id or (existing or {}).get('finalization_job_id'), 'dispatch_generation': dispatch_generation or (existing or {}).get('dispatch_generation'), } - self.pending_request_event.set() async def request_conversation_processing( self, @@ -138,7 +316,8 @@ async def request_conversation_processing( dispatch_generation: Optional[int] = None, ): """Request pusher to process a conversation through its durable lease.""" - if not self.pusher_connected or not self.pusher_ws: + pusher_ws = self.pusher_ws + if not self.pusher_connected or not pusher_ws: logger.info( f"Pusher not connected for {conversation_id}, will retry on reconnect {self.uid} {self.session_id}" ) @@ -166,15 +345,36 @@ async def request_conversation_processing( payload['finalization_job_id'] = pending['finalization_job_id'] payload['dispatch_generation'] = pending.get('dispatch_generation') or 1 data.extend(bytes(json.dumps(payload), "utf-8")) - await self.pusher_ws.send(cast(bytes, data)) + await pusher_ws.send(cast(bytes, data)) logger.info(f"Sent process_conversation request to pusher: {conversation_id} {self.uid} {self.session_id}") return True + except asyncio.CancelledError: + self._mark_failed_socket_disconnected( + pusher_ws, + auto_reconnect=self.deps.is_active(), + ) + raise except Exception as e: logger.error(f"Failed to send process_conversation request: {e} {self.uid} {self.session_id}") + self._mark_failed_socket_disconnected( + pusher_ws, + auto_reconnect=self.deps.is_active(), + ) return False async def _transcript_flush(self, auto_reconnect: bool = True): - if self.pusher_connected and self.pusher_ws and len(self.segment_buffers) > 0: + async with self.transcript_flush_lock: + pusher_ws = self.pusher_ws + if not self.pusher_connected or not pusher_ws: + return + + delivery = self._take_transcript_delivery() + if delivery is None: + return + if self.delivery_ack_supported and not self._delivery_due( + delivery.last_sent_ws, delivery.last_sent_at, pusher_ws + ): + return try: data = bytearray() data.extend(struct.pack("I", 102)) @@ -182,106 +382,160 @@ async def _transcript_flush(self, auto_reconnect: bool = True): bytes( json.dumps( { - "segments": list(self.segment_buffers), - "memory_id": self.deps.get_current_conversation_id(), + "segments": delivery.segments, + "memory_id": delivery.conversation_id, + "delivery_id": delivery.delivery_id, } ), "utf-8", ) ) - self.segment_buffers.clear() - await self.pusher_ws.send(cast(bytes, data)) + except Exception as e: + logger.error(f"Pusher transcripts serialization failed: {e} {self.uid} {self.session_id}") + return + + try: + await pusher_ws.send(cast(bytes, data)) + if self.delivery_ack_supported: + delivery.last_sent_ws = pusher_ws + delivery.last_sent_at = self.deps.monotonic() + elif self.pending_transcript_delivery is delivery: + self.pending_transcript_delivery = None + except asyncio.CancelledError: + # The stable pending delivery remains available to a reconnecting + # session even though the caller is being cancelled. + self._mark_failed_socket_disconnected( + pusher_ws, + auto_reconnect=self.deps.is_active(), + ) + raise except ConnectionClosed as e: logger.error(f"Pusher transcripts Connection closed: {e} {self.uid} {self.session_id}") - self._mark_disconnected() + self._mark_failed_socket_disconnected(pusher_ws, auto_reconnect=auto_reconnect) except Exception as e: logger.error(f"Pusher transcripts failed: {e} {self.uid} {self.session_id}") + self._mark_failed_socket_disconnected(pusher_ws, auto_reconnect=auto_reconnect) async def transcript_consume(self): while self.deps.is_active(): await self.deps.sleep(1) - if len(self.segment_buffers) > 0: + if self.pending_transcript_delivery is not None or self.segment_buffers: await self._transcript_flush(auto_reconnect=True) def audio_bytes_send(self, audio_bytes: bytes, received_at: float): + max_size = max(0, self.config.max_audio_buffer_size) + pending_size = ( + sum(len(chunk) for chunk in self.pending_audio_delivery.chunks) if self.pending_audio_delivery else 0 + ) + live_capacity = max(0, max_size - pending_size) + if live_capacity == 0: + record_fallback( + component='pusher', + from_mode='listen_audio_buffer', + to_mode='drop_newest', + reason='capacity_full', + outcome='degraded', + log=logger, + ) + return chunk = audio_bytes - if len(chunk) > self.config.max_audio_buffer_size: - chunk = chunk[-self.config.max_audio_buffer_size :] - while self.audio_total_size + len(chunk) > self.config.max_audio_buffer_size and self.audio_chunks: + dropped = False + if len(chunk) > live_capacity: + chunk = chunk[-live_capacity:] + dropped = True + while self.audio_total_size + len(chunk) > live_capacity and self.audio_chunks: old = self.audio_chunks.popleft() + self.audio_chunk_conversation_ids.popleft() + self.audio_chunk_received_at.popleft() self.audio_total_size -= len(old) + dropped = True self.audio_chunks.append(chunk) + self.audio_chunk_conversation_ids.append(self.deps.get_current_conversation_id()) + self.audio_chunk_received_at.append(received_at) self.audio_total_size += len(chunk) self.audio_buffer_last_received = received_at + if dropped: + record_fallback( + component='pusher', + from_mode='listen_audio_buffer', + to_mode='drop_oldest', + reason='capacity_full', + outcome='degraded', + log=logger, + ) async def _audio_bytes_flush(self, auto_reconnect: bool = True): - current_conversation_id = self.deps.get_current_conversation_id() - if ( - self.pusher_ws - and current_conversation_id - and ( - self.last_synced_conversation_id is None or current_conversation_id != self.last_synced_conversation_id - ) - ): - try: - data = bytearray() - data.extend(struct.pack("I", 103)) - data.extend(bytes(current_conversation_id, "utf-8")) - await self.pusher_ws.send(cast(bytes, data)) - self.last_synced_conversation_id = current_conversation_id - except ConnectionClosed as e: - logger.error(f"Pusher audio_bytes Connection closed: {e} {self.uid} {self.session_id}") - self._mark_disconnected() - except Exception as e: - logger.error(f"Failed to send conversation_id to pusher: {e} {self.uid} {self.session_id}") + async with self.audio_flush_lock: + pusher_ws = self.pusher_ws + if not self.pusher_connected or not pusher_ws: + return + + delivery = self._take_audio_delivery() + if delivery is None: + return - if self.pusher_connected and self.pusher_ws and self.audio_total_size > 0: try: + if delivery.conversation_id and delivery.conversation_id != self.last_synced_conversation_id: + data = bytearray() + data.extend(struct.pack("I", 103)) + data.extend(bytes(delivery.conversation_id, "utf-8")) + await pusher_ws.send(cast(bytes, data)) + if self.pusher_ws is pusher_ws: + self.last_synced_conversation_id = delivery.conversation_id + + audio_data = b''.join(delivery.chunks) effective_rate = TARGET_SAMPLE_RATE if self.config.is_multi_channel else self.config.sample_rate - buffer_duration_seconds = self.audio_total_size / (effective_rate * 2) - buffer_start_time = (self.audio_buffer_last_received or self.deps.now()) - buffer_duration_seconds - audio_data = b''.join(self.audio_chunks) + buffer_duration_seconds = len(audio_data) / (effective_rate * 2) + buffer_start_time = (delivery.last_received or self.deps.now()) - buffer_duration_seconds data = bytearray() data.extend(struct.pack("I", 101)) data.extend(struct.pack("d", buffer_start_time)) data.extend(audio_data) - self.audio_chunks.clear() - self.audio_total_size = 0 - del audio_data - await self.pusher_ws.send(cast(bytes, data)) + await pusher_ws.send(cast(bytes, data)) + if self.pending_audio_delivery is delivery: + self.pending_audio_delivery = None + except asyncio.CancelledError: + # Keep the route-stamped delivery intact for replay, then honor + # structured cancellation. + self._mark_failed_socket_disconnected( + pusher_ws, + auto_reconnect=self.deps.is_active(), + ) + raise except ConnectionClosed as e: logger.error(f"Pusher audio_bytes Connection closed: {e} {self.uid} {self.session_id}") - self._mark_disconnected() + self._mark_failed_socket_disconnected(pusher_ws, auto_reconnect=auto_reconnect) except Exception as e: logger.error(f"Pusher audio_bytes failed: {e} {self.uid} {self.session_id}") + self._mark_failed_socket_disconnected(pusher_ws, auto_reconnect=auto_reconnect) async def audio_bytes_consume(self): while self.deps.is_active(): await self.deps.sleep(1) - if self.audio_total_size > 0: + if self.pending_audio_delivery is not None or self.audio_total_size > 0: await self._audio_bytes_flush(auto_reconnect=True) async def pusher_receive(self): """Receive and handle messages from pusher, with timeout-based retry for pending requests.""" while self.deps.is_active(): - if not self.pending_conversation_requests: - self.pending_request_event.clear() - try: - await asyncio.wait_for(self.pending_request_event.wait(), timeout=5.0) - except asyncio.TimeoutError: - continue - - if not self.pusher_connected or not self.pusher_ws: + pusher_ws = self.pusher_ws + if not self.pusher_connected or pusher_ws is None: await self.deps.sleep(0.5) continue try: - msg = cast(bytes, await asyncio.wait_for(self.pusher_ws.recv(), timeout=5.0)) + async with self.pusher_receive_lock: + msg = cast(bytes, await asyncio.wait_for(pusher_ws.recv(), timeout=5.0)) if not msg or len(msg) < 4: continue header_type = struct.unpack(' None: + if not self.delivery_ack_supported or not self.pusher_connected or not self.pusher_ws: + return + + while self.pusher_connected and self.pusher_ws: + remaining = deadline - asyncio.get_running_loop().time() + if remaining <= 0: + return + if self.pending_transcript_delivery is not None or self.segment_buffers: + await asyncio.wait_for( + self._transcript_flush(auto_reconnect=False), + timeout=remaining, + ) + remaining = deadline - asyncio.get_running_loop().time() + if remaining <= 0: + return + await asyncio.wait_for( + self._retry_pending_speaker_sample_requests(), + timeout=remaining, + ) + if ( + self.pending_transcript_delivery is None + and not self.segment_buffers + and not self.pending_speaker_sample_requests + ): + return + + remaining = deadline - asyncio.get_running_loop().time() + if remaining <= 0: + logger.warning( + 'Pusher close acknowledgement deadline elapsed pending_transcript=%s buffered_transcript=%s pending_speaker=%s uid=%s session=%s', + self.pending_transcript_delivery is not None, + len(self.segment_buffers), + len(self.pending_speaker_sample_requests), + self.uid, + self.session_id, + ) + return + try: + pusher_ws = self.pusher_ws + + async def receive_one() -> bytes: + async with self.pusher_receive_lock: + return cast(bytes, await pusher_ws.recv()) + + msg = await asyncio.wait_for( + receive_one(), + timeout=min(remaining, PUSHER_DELIVERY_ACK_TIMEOUT), + ) + except (asyncio.TimeoutError, ConnectionClosed): + if asyncio.get_running_loop().time() >= deadline: + return + continue + if not msg or len(msg) < 4 or struct.unpack(' None: + if not self.delivery_ack_supported or not self.pusher_connected or not self.pusher_ws: + return + await self.pusher_ws.send(struct.pack('I', PUSHER_DELIVERY_DRAIN_OPCODE)) + async def close(self, code: int = 1000): if self.reconnect_task and not self.reconnect_task.done(): self.reconnect_task.cancel() @@ -519,9 +878,38 @@ async def close(self, code: int = 1000): except asyncio.CancelledError: pass self.reconnect_task = None - await self._flush() - if self.pusher_ws: - await self.pusher_ws.close(code) + deadline = asyncio.get_running_loop().time() + PUSHER_CLOSE_ACK_TIMEOUT + + async def flush_and_drain() -> None: + await self._flush() + await self._retry_pending_speaker_sample_requests() + await self._request_delivery_drain() + await self._drain_delivery_acks(deadline) + + try: + await asyncio.wait_for( + flush_and_drain(), + timeout=max(0.0, deadline - asyncio.get_running_loop().time()), + ) + except asyncio.TimeoutError: + logger.warning( + 'Pusher close delivery deadline elapsed pending_transcript=%s buffered_transcript=%s pending_speaker=%s uid=%s session=%s', + self.pending_transcript_delivery is not None, + len(self.segment_buffers), + len(self.pending_speaker_sample_requests), + self.uid, + self.session_id, + ) + finally: + pusher_ws = self.pusher_ws + if pusher_ws is not None: + try: + await asyncio.wait_for( + pusher_ws.close(code), + timeout=PUSHER_SOCKET_CLOSE_TIMEOUT, + ) + except asyncio.TimeoutError: + logger.warning('Pusher socket close deadline elapsed uid=%s session=%s', self.uid, self.session_id) def is_degraded(self): return self.reconnect_state in (PusherReconnectState.DEGRADED, PusherReconnectState.HALF_OPEN_PROBE) @@ -533,12 +921,38 @@ async def send_speaker_sample_request( segment_ids: List[str], ): """Send speaker sample extraction request to pusher with segment IDs.""" - if not self.pusher_connected or not self.pusher_ws: - self.pending_speaker_sample_requests.append((person_id, conv_id, segment_ids)) - logger.warning( - f"Pusher not connected, buffered speaker sample request: person={person_id}, " - f"{len(segment_ids)} segments ({len(self.pending_speaker_sample_requests)} pending) {self.uid} {self.session_id}" - ) + async with self.speaker_sample_send_lock: + if not self._buffer_pending_speaker_sample_request(person_id, conv_id, segment_ids): + return + if not self.pusher_connected or not self.pusher_ws: + logger.warning( + f"Pusher not connected, buffered speaker sample request: person={person_id}, " + f"{len(segment_ids)} segments ({len(self.pending_speaker_sample_requests)} pending) {self.uid} {self.session_id}" + ) + return + await self._send_speaker_sample_request(person_id, conv_id, segment_ids) + + async def _send_speaker_sample_request( + self, + person_id: str, + conv_id: str, + segment_ids: List[str], + ) -> None: + pusher_ws = self.pusher_ws + if not self.pusher_connected or not pusher_ws: + return + + request = (person_id, conv_id, list(segment_ids)) + request_key = (person_id, conv_id, tuple(segment_ids)) + delivery_id = self.pending_speaker_sample_delivery_ids.get(request_key) + if delivery_id is None: + if not self._buffer_pending_speaker_sample_request(person_id, conv_id, segment_ids): + return + delivery_id = self.pending_speaker_sample_delivery_ids.get(request_key) + if delivery_id is None: + return + sent = self.pending_speaker_sample_sent.get(request_key) + if self.delivery_ack_supported and sent is not None and not self._delivery_due(sent[0], sent[1], pusher_ws): return try: data = bytearray() @@ -550,17 +964,35 @@ async def send_speaker_sample_request( "person_id": person_id, "conversation_id": conv_id, "segment_ids": segment_ids, + "delivery_id": delivery_id, } ), "utf-8", ) ) - await self.pusher_ws.send(cast(bytes, data)) + await pusher_ws.send(cast(bytes, data)) + if self.delivery_ack_supported: + self.pending_speaker_sample_sent[request_key] = (pusher_ws, self.deps.monotonic()) + else: + try: + self.pending_speaker_sample_requests.remove(request) + except ValueError: + pass + self.pending_speaker_sample_delivery_ids.pop(request_key, None) + self.pending_speaker_sample_sent.pop(request_key, None) logger.info( f"Sent speaker sample request to pusher: person={person_id}, {len(segment_ids)} segments {self.uid} {self.session_id}" ) + except asyncio.CancelledError: + # Request and delivery id remain buffered for a safe replay. + self._mark_failed_socket_disconnected( + pusher_ws, + auto_reconnect=self.deps.is_active(), + ) + raise except Exception as e: logger.error(f"Failed to send speaker sample request: {e} {self.uid} {self.session_id}") + self._mark_failed_socket_disconnected(pusher_ws, auto_reconnect=True) def is_connected(self): return self.pusher_connected @@ -570,11 +1002,12 @@ async def pusher_heartbeat(self): while self.deps.is_active(): if await self.deps.wait_for_event(self.deps.shutdown_event, 20): break - if self.pusher_connected and self.pusher_ws: + pusher_ws = self.pusher_ws + if self.pusher_connected and pusher_ws: try: - await self.pusher_ws.send(struct.pack("I", 100)) + await pusher_ws.send(struct.pack("I", 100)) except ConnectionClosed: - self._mark_disconnected() + self._mark_failed_socket_disconnected(pusher_ws, auto_reconnect=True) except Exception as e: logger.error(f"Pusher heartbeat send failed: {e} {self.uid} {self.session_id}") diff --git a/backend/utils/memory/memory_service.py b/backend/utils/memory/memory_service.py index 97cd760e8fb..9e4f6878e42 100644 --- a/backend/utils/memory/memory_service.py +++ b/backend/utils/memory/memory_service.py @@ -10,7 +10,11 @@ import database.memories as memories_db import database.vector_db as vector_db -from database.vector_db import delete_memory_vector, upsert_memory_vector, upsert_memory_vectors_batch +from database.vector_db import ( + delete_memory_vector, + upsert_memory_vector, + upsert_memory_vectors_batch, +) from models.memories import MemoryDB from utils.memory.canonical_memory_adapter import ( delete_all_canonical_memories, @@ -39,6 +43,7 @@ MemoryPayload = Dict[str, Any] McpSearchPayload = Dict[str, Any] +LEGACY_SEARCH_MAX_CANDIDATES = 60 class DeviceScopeNotSupportedError(ValueError): @@ -183,74 +188,137 @@ def _memory_ids_and_scores(matches: List[MemoryPayload]) -> tuple[List[str], Dic memory_id = match.get("memory_id") if not isinstance(memory_id, str) or not memory_id: continue + if memory_id in scores_by_id: + continue memory_ids.append(memory_id) scores_by_id[memory_id] = float(match.get("score") or 0) return memory_ids, scores_by_id +def _legacy_search_fetch_limits(capped_limit: int) -> List[int]: + """Grow Pinecone top_k only when filtering leaves a result page underfilled.""" + fetch_limit = min(capped_limit * 3, LEGACY_SEARCH_MAX_CANDIDATES) + limits = [fetch_limit] + while fetch_limit < LEGACY_SEARCH_MAX_CANDIDATES: + fetch_limit = min(fetch_limit * 2, LEGACY_SEARCH_MAX_CANDIDATES) + limits.append(fetch_limit) + return limits + + +def _hydrate_legacy_search_memories( + uid: str, + memory_ids: List[str], + *, + requested_memory_ids: set[str], + memories_by_id: Dict[str, MemoryPayload], +) -> None: + new_memory_ids = [memory_id for memory_id in memory_ids if memory_id not in requested_memory_ids] + if not new_memory_ids: + return + requested_memory_ids.update(new_memory_ids) + for memory in memories_db.get_memories_by_ids(uid, new_memory_ids): + memory_id = memory.get("id") + if isinstance(memory_id, str) and memory_id: + memories_by_id[memory_id] = memory + + def _legacy_search_memories(uid: str, query: str, *, limit: int = 5) -> List[MemorySearchMatch]: capped_limit = max(1, min(limit, 20)) - matches = vector_db.find_similar_memories(uid, query, threshold=0.0, limit=capped_limit) - if not matches: - return [] - - memory_ids, scores_by_id = _memory_ids_and_scores(matches) - if not memory_ids: - return [] - - memories_data = memories_db.get_memories_by_ids(uid, memory_ids) - memories_data = [ - memory_api_payload(memory, MemoryApiExposure.LEGACY) - for memory in memories_data - if not memory.get("is_locked", False) - ] - + memories_by_id: Dict[str, MemoryPayload] = {} + requested_memory_ids: set[str] = set() results: List[MemorySearchMatch] = [] - for memory_data in memories_data: - memory_id = memory_data.get("id") - if not isinstance(memory_id, str): - continue - try: - memory_obj = _legacy_memorydb(memory_data) - except ValidationError: - continue - results.append(MemorySearchMatch(memory=memory_obj, score=scores_by_id.get(memory_id, 0.0))) + for fetch_limit in _legacy_search_fetch_limits(capped_limit): + matches = vector_db.find_similar_memories(uid, query, threshold=0.0, limit=fetch_limit) + if not matches: + break + + memory_ids, scores_by_id = _memory_ids_and_scores(matches) + _hydrate_legacy_search_memories( + uid, + memory_ids, + requested_memory_ids=requested_memory_ids, + memories_by_id=memories_by_id, + ) + + results = [] + for memory_id in memory_ids: + memory_data = memories_by_id.get(memory_id) + if not memory_data: + continue + if ( + memory_data.get("is_locked", False) + or memory_data.get("user_review") is False + or memory_data.get("invalid_at") is not None + ): + continue + try: + memory_obj = _legacy_memorydb(memory_api_payload(memory_data, MemoryApiExposure.LEGACY)) + except ValidationError: + continue + results.append(MemorySearchMatch(memory=memory_obj, score=scores_by_id.get(memory_id, 0.0))) + if len(results) >= capped_limit: + return results + if len(matches) < fetch_limit: + break return results +def search_legacy_memories(uid: str, query: str, *, limit: int = 5) -> List[MemorySearchMatch]: + """Shared filtered legacy search for API and chat-tool surfaces.""" + return _legacy_search_memories(uid, query, limit=limit) + + +def _single_vector_write_succeeded(result: Any) -> bool: + if result is None or result is False: + return False + if isinstance(result, (list, tuple)) and not result: + return False + return True + + +def _batch_vector_write_succeeded(result: Any, expected_count: int) -> bool: + return type(result) is int and result == expected_count + + def _legacy_search_memories_mcp(uid: str, query: str, *, limit: int = 5) -> List[McpSearchPayload]: """Legacy MCP search path: over-fetch, filter, RRF rerank (Wave 2 cf#1 parity).""" capped_limit = max(1, min(limit, 20)) - fetch_limit = min(capped_limit * 3, 60) - matches = vector_db.find_similar_memories(uid, query, threshold=0.0, limit=fetch_limit) - if not matches: - return [] - - memory_ids, scores = _memory_ids_and_scores(matches) - if not memory_ids: - return [] - docs: Dict[str, MemoryPayload] = {} - for memory in memories_db.get_memories_by_ids(uid, memory_ids): - memory_id = memory.get("id") - if isinstance(memory_id, str) and memory_id: - docs[memory_id] = memory - + requested_memory_ids: set[str] = set() candidates: List[McpSearchPayload] = [] - for memory_id in memory_ids: - memory = docs.get(memory_id) - if not memory: - continue - if memory.get("user_review") is False or memory.get("is_locked", False) or memory.get("invalid_at") is not None: - continue - candidates.append( - { - "id": memory.get("id", ""), - "content": memory.get("content", ""), - "category": memory.get("category", "other"), - "vector_score": scores.get(memory_id, 0), - } + for fetch_limit in _legacy_search_fetch_limits(capped_limit): + matches = vector_db.find_similar_memories(uid, query, threshold=0.0, limit=fetch_limit) + if not matches: + break + + memory_ids, scores = _memory_ids_and_scores(matches) + _hydrate_legacy_search_memories( + uid, + memory_ids, + requested_memory_ids=requested_memory_ids, + memories_by_id=docs, ) + candidates = [] + for memory_id in memory_ids: + memory = docs.get(memory_id) + if not memory: + continue + if ( + memory.get("user_review") is False + or memory.get("is_locked", False) + or memory.get("invalid_at") is not None + ): + continue + candidates.append( + { + "id": memory.get("id", ""), + "content": memory.get("content", ""), + "category": memory.get("category", "other"), + "vector_score": scores.get(memory_id, 0), + } + ) + if len(candidates) >= capped_limit or len(matches) < fetch_limit: + break candidates.sort(key=lambda candidate: candidate.get("vector_score", 0), reverse=True) reranked = rrf_rerank(query, candidates, capped_limit) @@ -610,16 +678,26 @@ def create_external_memory( raise HTTPException(status_code=503, detail="Service temporarily unavailable") _require_legacy_write_guard(uid, self._db_client, consumer=consumer, operation=operation) - memories_db.create_memory(uid, memory_write_payload(memory_db, MemoryApiExposure.LEGACY)) + memories_db.create_memory( + uid, + memory_write_payload(memory_db, MemoryApiExposure.LEGACY), + firestore_client=self._db_client, + ) if upsert_vector: try: - upsert_memory_vector( + projection_result = upsert_memory_vector( uid, memory_db.id, memory_db.content, memory_db.category.value, subject_entity_id=memory_db.subject_entity_id, ) + if not _single_vector_write_succeeded(projection_result): + logger.warning( + "Vector upsert returned no write uid=%s memory_id=%s (memory saved, vector missing)", + uid, + memory_db.id, + ) except Exception: logger.exception( "Vector upsert failed uid=%s memory_id=%s (memory saved, vector missing)", @@ -662,10 +740,11 @@ def create_external_memory_batch( memories_db.save_memories( uid, [memory_write_payload(memory, MemoryApiExposure.LEGACY) for memory in memory_dbs], + firestore_client=self._db_client, ) if upsert_vectors: try: - upsert_memory_vectors_batch( + projection_result = upsert_memory_vectors_batch( uid, [ { @@ -677,6 +756,13 @@ def create_external_memory_batch( for memory in memory_dbs ], ) + if not _batch_vector_write_succeeded(projection_result, len(memory_dbs)): + logger.warning( + "Vector batch upsert returned partial/no write uid=%s expected=%d actual=%r", + uid, + len(memory_dbs), + projection_result, + ) except Exception: logger.exception("Vector batch upsert failed uid=%s (memories saved, vectors missing)", uid) return [_legacy_memorydb(memory) for memory in memory_dbs] @@ -702,15 +788,29 @@ def delete_external_memory( return _require_legacy_write_guard(uid, self._db_client, consumer=consumer, operation=operation) - memory = memories_db.get_memory(uid, memory_id) + memory = memories_db.get_memory(uid, memory_id, firestore_client=self._db_client) if not memory: raise HTTPException(status_code=404, detail="Memory not found") if memory.get('is_locked', False): raise HTTPException(status_code=402, detail="A paid plan is required to access this memory.") - memories_db.delete_memory(uid, memory_id) + delete_result = memories_db.delete_memory(uid, memory_id, firestore_client=self._db_client) + if delete_result.committed_count != 1: + logger.error( + "Firestore delete count mismatch uid=%s memory_id=%s actual=%r", + uid, + memory_id, + delete_result.committed_count, + ) + raise HTTPException(status_code=503, detail="Service temporarily unavailable") if delete_vector: try: - delete_memory_vector(uid, memory_id) + projection_result = delete_memory_vector(uid, memory_id) + if projection_result is not True: + logger.warning( + "Vector delete returned no write uid=%s memory_id=%s (Firestore deleted)", + uid, + memory_id, + ) except Exception: logger.exception("Vector delete failed uid=%s memory_id=%s (Firestore deleted)", uid, memory_id) @@ -735,25 +835,41 @@ def update_external_memory_content( raise HTTPException(status_code=404, detail="Memory not found") _require_legacy_write_guard(uid, self._db_client, consumer=consumer, operation=operation) - memory = memories_db.get_memory(uid, memory_id) + memory = memories_db.get_memory(uid, memory_id, firestore_client=self._db_client) if not memory: raise HTTPException(status_code=404, detail="Memory not found") if memory.get('is_locked', False): raise HTTPException(status_code=402, detail="A paid plan is required to access this memory.") - memories_db.edit_memory(uid, memory_id, content) + memories_db.edit_memory( + uid, + memory_id, + content, + firestore_client=self._db_client, + ) if upsert_vector: try: - upsert_memory_vector( + projection_result = upsert_memory_vector( uid, memory_id, content, memory.get('category', 'other'), subject_entity_id=memory.get('subject_entity_id'), ) + if not _single_vector_write_succeeded(projection_result): + logger.warning( + "Vector upsert returned no write uid=%s memory_id=%s (memory edited, vector stale)", + uid, + memory_id, + ) except Exception: logger.exception( "Vector upsert failed uid=%s memory_id=%s (memory edited, vector stale)", uid, memory_id, ) - return _legacy_memorydb(cast(MemoryPayload, memories_db.get_memory(uid, memory_id))) + return _legacy_memorydb( + cast( + MemoryPayload, + memories_db.get_memory(uid, memory_id, firestore_client=self._db_client), + ) + ) diff --git a/backend/utils/other/storage.py b/backend/utils/other/storage.py index b80e1241b74..bc9229d28ca 100644 --- a/backend/utils/other/storage.py +++ b/backend/utils/other/storage.py @@ -214,6 +214,7 @@ def upload_person_speech_sample_from_bytes( uid: str, person_id: str, sample_rate: int = 16000, + deduplication_key: Optional[str] = None, ) -> str: """Upload PCM audio bytes as WAV speech sample. Returns GCS path.""" import uuid as uuid_module @@ -224,14 +225,13 @@ def upload_person_speech_sample_from_bytes( wav_file.setsampwidth(2) # 16-bit audio wav_file.setframerate(sample_rate) wav_file.writeframes(audio_bytes) - bucket = _get_speech_profiles_bucket(required=True) assert bucket is not None # required=True raises if missing - filename = f"{uuid_module.uuid4()}.wav" + filename_id = hashlib.sha256(deduplication_key.encode()).hexdigest() if deduplication_key else uuid_module.uuid4() + filename = f"{filename_id}.wav" path = f'{uid}/people_profiles/{person_id}/{filename}' blob = bucket.blob(path) blob.upload_from_string(wav_buffer.getvalue(), content_type='audio/wav') - return path diff --git a/backend/utils/retrieval/tool_services/memories.py b/backend/utils/retrieval/tool_services/memories.py index 00d12fbd6df..f0c7c789b25 100644 --- a/backend/utils/retrieval/tool_services/memories.py +++ b/backend/utils/retrieval/tool_services/memories.py @@ -4,15 +4,14 @@ """ from datetime import datetime, timezone -from typing import Optional, Any, Dict, List, cast +from typing import Optional, Any, Dict, List import database.memories as memory_db import database.notifications as notification_db -import database.vector_db as vector_db from database._client import db as firestore_db from models.memories import MemoryDB from utils.conversations.render import format_local_date, resolve_display_tz -from utils.memory.memory_service import MemoryService +from utils.memory.memory_service import MemoryService, search_legacy_memories from utils.memory.memory_system import MemorySystem from utils.memory.surface_routing import pin_memory_system from utils.memory.chat_memory_adapter import ( @@ -179,45 +178,17 @@ def search_memories_text( return default_memories.text or "No memories available for this request." try: - matches = vector_db.find_similar_memories(uid, query, threshold=0.0, limit=limit) - + matches = search_legacy_memories(uid, query, limit=limit) if not matches: return f"No memories found matching '{query}'." - memory_ids = [cast(str, match.get('memory_id')) for match in matches if match.get('memory_id')] - scores_by_id = {match.get('memory_id'): match.get('score', 0) for match in matches} - - if not memory_ids: - return f"Found matches but no valid memory IDs for query: '{query}'" - - memories_data = memory_db.get_memories_by_ids(uid, memory_ids) - - # Filter locked - memories_data = [m for m in memories_data if not m.get('is_locked', False)] - if not memories_data: - return f"No memories found matching '{query}'." - - # Format with scores - memory_objects: List[Dict[str, Any]] = [] - for memory_data in memories_data: - try: - memory_obj = MemoryDB(**memory_data) - score = scores_by_id.get(memory_data.get('id'), 0) - memory_objects.append({'memory': memory_obj, 'score': score}) - except Exception as e: - logger.error(f"Error creating MemoryDB object: {e}") - continue - - if not memory_objects: - return f"Found matches but could not retrieve memory details for query: '{query}'" - - result = f"Found {len(memory_objects)} memories matching '{query}':\n\n" - for item in memory_objects: - memory = item['memory'] - score = item['score'] + result = f"Found {len(matches)} memories matching '{query}':\n\n" + for match in matches: + memory = match.memory date_str = format_local_date(memory.created_at, display_tz) if memory.created_at else 'Unknown' result += ( - f"- {memory.content} (relevance: {score:.2f}, category: {memory.category.value}, date: {date_str})\n" + f"- {memory.content} (relevance: {match.score:.2f}, " + f"category: {memory.category.value}, date: {date_str})\n" ) return result.strip() diff --git a/backend/utils/speaker_identification.py b/backend/utils/speaker_identification.py index ac464a2fa24..9659a26ffbc 100644 --- a/backend/utils/speaker_identification.py +++ b/backend/utils/speaker_identification.py @@ -1,7 +1,8 @@ import io import re import wave -from typing import Any, Dict, List, Optional, cast +from dataclasses import dataclass +from typing import Any, Dict, List, Literal, Optional, cast import av import numpy as np @@ -21,6 +22,16 @@ logger = logging.getLogger(__name__) +@dataclass(frozen=True) +class SpeakerSampleExtractionResult: + status: Literal['stored', 'already_present', 'terminal_no_sample', 'retryable'] + reason: str + + @property + def retryable(self) -> bool: + return self.status == 'retryable' + + def _pcm_to_wav_bytes(pcm_data: bytes, sample_rate: int) -> bytes: """ Convert PCM16 mono audio to WAV format bytes. @@ -327,7 +338,8 @@ async def extract_speaker_samples( conversation_id: str, segment_ids: List[str], sample_rate: int = 16000, -): + delivery_id: Optional[str] = None, +) -> SpeakerSampleExtractionResult: """ Extract speech samples from segments and store as speaker profiles. Fetches conversation from DB to get started_at and segment details. @@ -337,6 +349,8 @@ async def extract_speaker_samples( # Run lazy migration for samples before checking count # (migration may drop invalid samples, freeing up space) person = await run_blocking(db_executor, users_db.get_person, uid, person_id) + if person is None: + return SpeakerSampleExtractionResult('retryable', 'person_not_ready') if person: person = await maybe_migrate_person_samples(uid, person) @@ -344,18 +358,18 @@ async def extract_speaker_samples( sample_count = await run_blocking(db_executor, users_db.get_person_speech_samples_count, uid, person_id) if sample_count >= 1: logger.warning(f"Person {person_id} already has {sample_count} samples, skipping {uid} {conversation_id}") - return + return SpeakerSampleExtractionResult('already_present', 'sample_already_present') # Fetch conversation to get started_at and segment details conversation = await run_blocking(db_executor, conversations_db.get_conversation, uid, conversation_id) if not conversation: logger.warning(f"Conversation {conversation_id} not found {uid}") - return + return SpeakerSampleExtractionResult('retryable', 'conversation_not_ready') started_at = conversation.get('started_at') if not started_at: logger.info(f"Conversation {conversation_id} has no started_at {uid}") - return + return SpeakerSampleExtractionResult('retryable', 'conversation_not_ready') started_at_ts = started_at.timestamp() if hasattr(started_at, 'timestamp') else float(started_at) @@ -367,7 +381,7 @@ async def extract_speaker_samples( audio_files = conversation.get('audio_files', []) if not audio_files: logger.warning(f"No audio files found for {conversation_id}, skipping speaker sample extraction {uid}") - return + return SpeakerSampleExtractionResult('retryable', 'audio_files_not_ready') # Collect all chunk timestamps from audio files all_timestamps: List[Any] = [] @@ -377,13 +391,14 @@ async def extract_speaker_samples( if not all_timestamps: logger.warning(f"No chunk timestamps found for {conversation_id}, skipping speaker sample extraction {uid}") - return + return SpeakerSampleExtractionResult('retryable', 'audio_timestamps_not_ready') # Build chunks list in expected format chunks: List[Dict[str, Any]] = [{'timestamp': ts} for ts in sorted(set(all_timestamps))] samples_added = 0 max_samples_to_add = 1 - sample_count + retryable_reason: Optional[str] = None # Build ordered list with index lookup for expansion ordered_segments = [s for s in conv_segments if s.get('id')] @@ -396,6 +411,7 @@ async def extract_speaker_samples( seg = segment_map.get(seg_id) if not seg: logger.warning(f"Segment {seg_id} not found in conversation {uid} {conversation_id}") + retryable_reason = 'transcript_segments_not_ready' continue segment_start = seg.get('start') @@ -464,6 +480,7 @@ async def extract_speaker_samples( logger.info( f"No relevant chunks for segment {segment_start:.1f}-{segment_end:.1f}s {uid} {conversation_id}" ) + retryable_reason = 'audio_chunks_not_ready' continue # Download, merge, and extract (sync_executor avoids parent-child deadlock on storage_executor, #7387) @@ -503,15 +520,30 @@ async def extract_speaker_samples( transcript, is_valid, reason = await verify_and_transcribe_sample(wav_bytes, sample_rate, expected_text) if not is_valid: logger.error(f"Sample failed quality check: {reason} {uid} {conversation_id}") + if reason.startswith('transcription_failed'): + retryable_reason = 'transcription_failed' continue # Try next segment # Upload and store + sample_deduplication_key = f'speaker-sample\0{uid}\0{person_id}\0{delivery_id}' if delivery_id else None path = await run_blocking( - storage_executor, upload_person_speech_sample_from_bytes, sample_audio, uid, person_id, sample_rate + storage_executor, + upload_person_speech_sample_from_bytes, + sample_audio, + uid, + person_id, + sample_rate, + sample_deduplication_key, ) success = await run_blocking( - db_executor, users_db.add_person_speech_sample, uid, person_id, path, transcript=transcript + db_executor, + users_db.add_person_speech_sample, + uid, + person_id, + path, + transcript=transcript, + max_samples=1, ) if success: samples_added += 1 @@ -533,9 +565,20 @@ async def extract_speaker_samples( ) except Exception as emb_err: logger.error(f"Failed to extract/store speaker embedding: {emb_err} {uid} {conversation_id}") + return SpeakerSampleExtractionResult('stored', 'sample_stored') else: logger.error(f"Failed to add speech sample for person {person_id} {uid} {conversation_id}") - break # Likely hit limit + current_count = await run_blocking( + db_executor, users_db.get_person_speech_samples_count, uid, person_id + ) + if current_count >= 1: + return SpeakerSampleExtractionResult('already_present', 'sample_added_concurrently') + return SpeakerSampleExtractionResult('retryable', 'sample_persistence_failed') except Exception as e: logger.error(f"Error extracting speaker samples: {e} {uid} {conversation_id}") + return SpeakerSampleExtractionResult('retryable', 'extraction_failed') + + if retryable_reason is not None: + return SpeakerSampleExtractionResult('retryable', retryable_reason) + return SpeakerSampleExtractionResult('terminal_no_sample', 'no_eligible_segment') diff --git a/backend/utils/task_integrations_ops.py b/backend/utils/task_integrations_ops.py index c4f731d1641..96b411ada31 100644 --- a/backend/utils/task_integrations_ops.py +++ b/backend/utils/task_integrations_ops.py @@ -28,6 +28,51 @@ http_client: Optional[httpx.AsyncClient] = None +def _provider_create_success(external_task_id: Any) -> dict: + task_id = str(external_task_id).strip() if external_task_id is not None else '' + if not task_id: + return { + 'success': False, + 'error': 'Provider response omitted task identity', + 'error_code': 'invalid_provider_response', + 'retryable': False, + 'ambiguous': True, + } + return {'success': True, 'external_task_id': task_id} + + +def _provider_create_http_failure(provider: str, status_code: int) -> dict: + if 200 <= status_code < 300: + return { + 'success': False, + 'error': f'{provider} response did not contain a completed task', + 'error_code': 'invalid_provider_response', + 'status_code': status_code, + 'retryable': False, + 'ambiguous': True, + } + ambiguous = status_code in {408, 425} or status_code >= 500 + return { + 'success': False, + 'error': f'{provider} API error: {status_code}', + 'error_code': 'api_error', + 'status_code': status_code, + 'retryable': status_code == 429, + 'ambiguous': ambiguous, + } + + +def _provider_create_transport_failure(error: httpx.TransportError) -> dict: + safe_before_send = isinstance(error, (httpx.PoolTimeout, httpx.ConnectTimeout, httpx.ConnectError)) + return { + 'success': False, + 'error': type(error).__name__, + 'error_code': 'transport_error', + 'retryable': safe_before_send, + 'ambiguous': not safe_before_send, + } + + def get_http_client() -> httpx.AsyncClient: """Get or create the HTTP client instance.""" global http_client @@ -187,6 +232,48 @@ async def ensure_valid_oauth_token( return integration +def _task_create_configuration_failure(app_key: str, integration: dict) -> Optional[dict]: + if app_key not in OAUTH_CONFIGS: + return { + 'success': False, + 'error': f'Unsupported integration: {app_key}', + 'error_code': 'unsupported', + 'retryable': False, + 'ambiguous': False, + } + if integration.get('connected') is False: + name = OAUTH_CONFIGS[app_key]['name'] + return { + 'success': False, + 'error': f'{name} token refresh failed', + 'error_code': 'token_refresh_failed', + 'retryable': False, + 'ambiguous': False, + } + if not integration.get('access_token'): + return { + 'success': False, + 'error': f'No access token for {app_key}', + 'error_code': 'no_access_token', + 'retryable': False, + 'ambiguous': False, + } + required_field = { + 'asana': ('workspace_gid', 'No workspace configured', 'no_workspace'), + 'google_tasks': ('default_list_id', 'No task list configured', 'no_list'), + 'clickup': ('list_id', 'No list configured', 'no_list'), + }.get(app_key) + if required_field and not integration.get(required_field[0]): + return { + 'success': False, + 'error': required_field[1], + 'error_code': required_field[2], + 'retryable': False, + 'ambiguous': False, + } + return None + + async def perform_request_with_token_retry( uid: str, app_key: str, @@ -224,7 +311,7 @@ async def create_task_internal( Returns: dict: {"success": bool, "external_task_id": str, "error": str, "error_code": str} """ - if app_key in ['google_tasks', 'asana']: + if app_key in {'google_tasks', 'asana'}: integration = await ensure_valid_oauth_token( uid, app_key, @@ -232,18 +319,13 @@ async def create_task_internal( refresh_if_missing_expires_at=(app_key == 'google_tasks'), client=client, ) - # Use `is False` so a missing key (None) falls through to access_token - # validation below instead of blocking valid tokens on legacy records. - if integration.get('connected') is False: - name = OAUTH_CONFIGS.get(app_key, {'name': app_key}).get('name', app_key) - return {"success": False, "error": f"{name} token refresh failed", "error_code": "token_refresh_failed"} - - access_token = integration.get('access_token') - if not access_token: - return {"success": False, "error": f"No access token for {app_key}", "error_code": "no_access_token"} + preflight_error = _task_create_configuration_failure(app_key, integration) + if preflight_error is not None: + return preflight_error try: client = client or get_http_client() + access_token = str(integration['access_token']) if app_key == 'todoist': body = {'content': title, 'priority': 2} @@ -260,7 +342,7 @@ async def create_task_internal( if response.status_code in [200, 201]: task_data = response.json() - return {"success": True, "external_task_id": str(task_data.get('id'))} + return _provider_create_success(task_data.get('id')) else: if response.status_code == 401: await run_blocking( @@ -270,20 +352,13 @@ async def create_task_internal( 'todoist', {'connected': False}, ) - return { - "success": False, - "error": f"Todoist API error: {response.status_code}", - "error_code": "api_error", - } + return _provider_create_http_failure('Todoist', response.status_code) elif app_key == 'asana': - workspace_gid = integration.get('workspace_gid') + workspace_gid = str(integration['workspace_gid']) project_gid = integration.get('project_gid') user_gid = integration.get('user_gid') - if not workspace_gid: - return {"success": False, "error": "No workspace configured", "error_code": "no_workspace"} - task_data = {'name': title, 'workspace': workspace_gid} if description: task_data['notes'] = description @@ -305,22 +380,21 @@ async def _asana_post(c, token): uid, app_key, integration, _asana_post, client=client ) if retry_err: - return {"success": False, "error": "Asana token refresh failed", "error_code": "token_refresh_failed"} + return { + "success": False, + "error": "Asana token refresh failed", + "error_code": "token_refresh_failed", + "retryable": False, + } if response.status_code in [200, 201]: result = response.json() - return {"success": True, "external_task_id": result.get('data', {}).get('gid')} + return _provider_create_success(result.get('data', {}).get('gid')) else: - return { - "success": False, - "error": f"Asana API error: {response.status_code}", - "error_code": "api_error", - } + return _provider_create_http_failure('Asana', response.status_code) elif app_key == 'google_tasks': - list_id = integration.get('default_list_id') - if not list_id: - return {"success": False, "error": "No task list configured", "error_code": "no_list"} + list_id = str(integration['default_list_id']) task_data = {'title': title} if description: @@ -343,22 +417,17 @@ async def _google_tasks_post(c, token): "success": False, "error": "Google Tasks token refresh failed", "error_code": "token_refresh_failed", + "retryable": False, } if response.status_code in [200, 201]: result = response.json() - return {"success": True, "external_task_id": result.get('id')} + return _provider_create_success(result.get('id')) else: - return { - "success": False, - "error": f"Google Tasks API error: {response.status_code}", - "error_code": "api_error", - } + return _provider_create_http_failure('Google Tasks', response.status_code) elif app_key == 'clickup': - list_id = integration.get('list_id') - if not list_id: - return {"success": False, "error": "No list configured", "error_code": "no_list"} + list_id = str(integration['list_id']) task_data: dict[str, Any] = {'name': title} if description: @@ -374,17 +443,27 @@ async def _google_tasks_post(c, token): if response.status_code in [200, 201]: result = response.json() - return {"success": True, "external_task_id": result.get('id')} + return _provider_create_success(result.get('id')) else: - return { - "success": False, - "error": f"ClickUp API error: {response.status_code}", - "error_code": "api_error", - } - + return _provider_create_http_failure('ClickUp', response.status_code) else: - return {"success": False, "error": f"Unsupported integration: {app_key}", "error_code": "unsupported"} + return { + 'success': False, + 'error': f'Unsupported integration: {app_key}', + 'error_code': 'unsupported', + 'retryable': False, + 'ambiguous': False, + } + except httpx.TransportError as e: + logger.error(f"Error creating task in {app_key}: {e}") + return _provider_create_transport_failure(e) except Exception as e: logger.error(f"Error creating task in {app_key}: {e}") - return {"success": False, "error": str(e)} + return { + "success": False, + "error": type(e).__name__, + "error_code": "internal_error", + "retryable": False, + "ambiguous": True, + } diff --git a/backend/utils/webhooks.py b/backend/utils/webhooks.py index def58f4a239..df8cc0152c5 100644 --- a/backend/utils/webhooks.py +++ b/backend/utils/webhooks.py @@ -6,6 +6,8 @@ from typing import List, Optional from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit +import httpx + from database.redis_db import ( get_user_webhook_db, user_webhook_status_db, @@ -13,20 +15,32 @@ enable_user_webhook_db, set_user_webhook_db, ) -from database.webhook_health import record_dev_webhook_failure, record_dev_webhook_success, _DEV_FAILURE_THRESHOLD +from database.webhook_health import ( + record_dev_webhook_failure, + record_dev_webhook_success, + reset_dev_webhook_health, + _DEV_FAILURE_THRESHOLD, +) from models.conversation import Conversation from models.users import WebhookType import database.notifications as notification_db from utils.conversations.render import populate_speaker_names, populate_folder_names from utils.conversations.render import conversation_to_dict from utils.executors import db_executor, run_blocking -from utils.http_client import get_webhook_client, get_webhook_circuit_breaker, get_webhook_semaphore +from utils.http_client import ( + get_webhook_client, + get_webhook_circuit_breaker, + get_webhook_semaphore, + reset_webhook_circuit_breaker, +) from utils.notifications import send_notification import logging logger = logging.getLogger(__name__) _DEV_WEBHOOK_RETRY_DELAYS = (1.0, 5.0, 30.0) +_DEV_WEBHOOK_RETRYABLE_STATUS_CODES = frozenset({408, 425, 429}) +_REALTIME_DEV_WEBHOOK_RETRY_DELAYS = (0.5, 2.0) def _get_dev_webhook_retry_delays() -> tuple[float, ...]: @@ -50,6 +64,20 @@ def _append_query_params(url: str, params: dict) -> str: return urlunsplit((parts.scheme, parts.netloc, parts.path, urlencode(query_items), parts.fragment)) +def _is_retryable_dev_webhook_status(status_code: int) -> bool: + return status_code in _DEV_WEBHOOK_RETRYABLE_STATUS_CODES or 500 <= status_code < 600 + + +def reset_user_webhook_delivery_health(uid: str, wtype: WebhookType, webhook_url: Optional[str]) -> None: + """Reset persisted and process-local failure gates after an explicit enable.""" + reset_dev_webhook_health(uid, wtype) + target_url = webhook_url or '' + if wtype == WebhookType.audio_bytes: + target_url = target_url.split(',', 1)[0] + if target_url: + reset_webhook_circuit_breaker(target_url) + + async def _post_dev_webhook( webhook_name: str, webhook_url: str, @@ -85,7 +113,13 @@ async def _post_dev_webhook( ) return response failure_reason = f'HTTP {response.status_code}' - except Exception as e: + if not _is_retryable_dev_webhook_status(response.status_code): + logger.error( + f'{webhook_name}: delivery failed status={response.status_code} ' + f'attempt={attempt_number}/{attempts} retryable=false' + ) + return response + except httpx.TransportError as e: last_response = None last_exception = e failure_reason = type(e).__name__ @@ -110,13 +144,13 @@ async def _post_dev_webhook( raise last_exception -async def _handle_dev_webhook_disable(uid: str, wtype: str, should_disable: bool): +async def _handle_dev_webhook_disable(uid: str, wtype: WebhookType | str, should_disable: bool): if should_disable: logger.warning( f'Dev webhook auto-disabled: uid={uid} type={wtype} after {_DEV_FAILURE_THRESHOLD} consecutive failures' ) await run_blocking(db_executor, disable_user_webhook_db, uid, wtype) - wtype_str = wtype.value if hasattr(wtype, 'value') else str(wtype) + wtype_str = wtype.value if isinstance(wtype, WebhookType) else str(wtype) await run_blocking( db_executor, send_notification, @@ -134,7 +168,7 @@ def _build_conversation_webhook_payload_sync(uid: str, memory: Conversation) -> return payload -async def conversation_created_webhook(uid, memory: Conversation): +async def conversation_created_webhook(uid: str, memory: Conversation): if memory.is_locked: return @@ -237,7 +271,12 @@ async def day_summary_webhook(uid, summary: str, summary_json: Optional[dict] = return -async def realtime_transcript_webhook(uid, segments: List[dict]): +async def realtime_transcript_webhook( + uid, + segments: List[dict], + *, + idempotency_key: Optional[str] = None, +): logger.info(f"realtime_transcript_webhook {uid}") toggled = await run_blocking(db_executor, user_webhook_status_db, uid, WebhookType.realtime_transcript) @@ -256,6 +295,8 @@ async def realtime_transcript_webhook(uid, segments: List[dict]): webhook_url, json={'segments': segments, 'session_id': uid}, headers={'Content-Type': 'application/json'}, + idempotency_key=idempotency_key, + retry_delays=_REALTIME_DEV_WEBHOOK_RETRY_DELAYS, ) if response.status_code >= 200 and response.status_code < 300: cb.record_success() @@ -331,6 +372,7 @@ async def send_audio_bytes_developer_webhook(uid: str, sample_rate: int, data: b webhook_url, content=bytes(data), headers={'Content-Type': 'application/octet-stream'}, + retry_delays=_REALTIME_DEV_WEBHOOK_RETRY_DELAYS, ) if response.status_code >= 200 and response.status_code < 300: cb.record_success() diff --git a/backend/utils/x_connector.py b/backend/utils/x_connector.py index 5cb69223175..b8bc1851f56 100644 --- a/backend/utils/x_connector.py +++ b/backend/utils/x_connector.py @@ -377,8 +377,12 @@ def _extract_and_index(uid: str, posts: List[Dict]) -> int: for mdb in memory_dbs: memory_service.write(uid, mdb.model_dump()) else: - memories_db.save_memories(uid, [memory_write_payload(m, MemoryApiExposure.LEGACY) for m in memory_dbs]) - upsert_memory_vectors_batch( + memories_db.save_memories( + uid, + [memory_write_payload(m, MemoryApiExposure.LEGACY) for m in memory_dbs], + firestore_client=db, + ) + projected_count = upsert_memory_vectors_batch( uid, [ { @@ -390,6 +394,10 @@ def _extract_and_index(uid: str, posts: List[Dict]) -> int: for m in memory_dbs ], ) + if projected_count != len(memory_dbs): + raise RuntimeError( + "X memory vector upsert was partial " f"expected={len(memory_dbs)} actual={projected_count!r}" + ) total += len(memory_dbs) # Do not acknowledge this raw source until every write above succeeds. diff --git a/desktop/windows/src/renderer/src/lib/omiApi.generated.ts b/desktop/windows/src/renderer/src/lib/omiApi.generated.ts index 0f6694e748b..1529de10ece 100644 --- a/desktop/windows/src/renderer/src/lib/omiApi.generated.ts +++ b/desktop/windows/src/renderer/src/lib/omiApi.generated.ts @@ -1224,8 +1224,11 @@ export interface CreateTaskRequest { } export interface CreateTaskResponse { + ambiguous?: boolean | null; error?: string | null; + error_code?: string | null; external_task_id?: string | null; + retryable?: boolean | null; success: boolean; } diff --git a/docs/api-reference/app-client-openapi.json b/docs/api-reference/app-client-openapi.json index 732f27a64a0..04b991a54a0 100644 --- a/docs/api-reference/app-client-openapi.json +++ b/docs/api-reference/app-client-openapi.json @@ -7765,6 +7765,17 @@ "CreateTaskResponse": { "description": "Response for task creation", "properties": { + "ambiguous": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "title": "Ambiguous" + }, "error": { "anyOf": [ { @@ -7776,6 +7787,17 @@ ], "title": "Error" }, + "error_code": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Error Code" + }, "external_task_id": { "anyOf": [ { @@ -7787,6 +7809,17 @@ ], "title": "External Task Id" }, + "retryable": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "title": "Retryable" + }, "success": { "title": "Success", "type": "boolean" diff --git a/docs/doc/developer/backend/listen_pusher_pipeline.mdx b/docs/doc/developer/backend/listen_pusher_pipeline.mdx index b419d97ea57..7854197f097 100644 --- a/docs/doc/developer/backend/listen_pusher_pipeline.mdx +++ b/docs/doc/developer/backend/listen_pusher_pipeline.mdx @@ -6,7 +6,7 @@ description: "Sequence diagrams for the /v4/listen WebSocket and Pusher processi # Listen + Pusher Pipeline — Sequence Diagrams -> Last updated: 2026-07-15 (durable recording-session bindings, fenced late-content recovery, and ordered lifecycle envelopes) +> Last updated: 2026-07-24 (route-stamped replay and effect-completion acknowledgements) > > These diagrams document the real behavior observed during E2E testing with > live services (backend, pusher, STT providers, embedding API). Update when the @@ -297,7 +297,41 @@ sequenceDiagram ## 3.1 Pusher Reconnect & Pending Flush -When pusher reconnects after a disconnection, all buffered conversations are replayed. +When pusher reconnects after a disconnection, all buffered conversations are replayed. Every new +pusher connection declares acknowledgement support with the +`X-Omi-Delivery-Ack: 1` WebSocket response header. Listen treats this capability as socket-local so +a rolling deploy can reconnect safely to either an old or new pusher pod. + +On an acknowledgement-capable socket, transcript and speaker-sample frames retain one stable +delivery ID after `send()` returns. Pusher rejects a stable frame without acknowledging it when the +bounded queue is full. A worker acquires a short Redis processing lease only when it is ready to own +the local invocation. After that invocation returns, pusher writes a seven-day done marker and +returns opcode `202`; a replay that finds the done marker skips the invocation and returns the same +acknowledgement. Speaker extraction reports not-yet-persisted conversation, transcript, or audio +metadata as retryable, so those outcomes release the lease without an acknowledgement. + +Realtime app and developer webhooks receive the stable delivery ID in their idempotency header. +Opcode `202` proves only that the pusher worker's local integration and webhook invocations +returned. It is not a downstream HTTP delivery receipt: those helpers retain their existing +best-effort retry, health, and circuit-breaker policies. The stable key lets a receiver deduplicate +repeated successful attempts, while the pusher protocol itself provides at-least-once local worker +invocation rather than exactly-once downstream delivery. A worker crash between an external effect +and the Redis done write can repeat that effect. Redis outages fail open only when a worker is ready +to process the frame and emit shared fallback telemetry. + +Graceful close sends opcode `107` after flushing tail audio. That asks pusher to flush private-cloud +audio metadata and bypass the speaker-sample age delay while the acknowledgement socket is still +open. Listen then performs a ten-second acknowledgement drain and gives socket close its own +two-second cap. Work still pending when that bounded drain ends is not durable across listen-process +teardown; stable replay applies while the session object remains alive or reconnecting. + +Audio keeps the rolling-compatible opcode `103` plus `101` format. Each buffered audio batch +captures its conversation route before a socket await, and an explicit send failure or cancellation +keeps that route-stamped batch for reconnect replay. Audio has no receiver acknowledgement or +cross-replica deduplication, so an ambiguous local send completion may lose or duplicate bytes. +Pending and live audio share one configured byte cap; pending and live transcripts likewise share +one segment cap. Transcript overflow rejects the newest segment instead of evicting previously +accepted FIFO work. Capacity drops emit fallback telemetry. ```mermaid sequenceDiagram @@ -409,6 +443,24 @@ Staleness is fingerprint-driven: when late chunks change `audio_files` (pusher b ## 6. Event Wire Protocol +### Backend to pusher binary frames + +| Opcode | Payload | Retry contract | +|--------|---------|----------------| +| `100` | Header only | Heartbeat | +| `101` | Start timestamp followed by PCM16LE audio | Best effort after the socket-local opcode `103` route; explicit send failures replay, but ambiguous completion has no receiver ACK | +| `102` | JSON transcript batch with `memory_id` and optional `delivery_id` | Stable delivery is retained until opcode `202` on negotiated sockets | +| `103` | UTF-8 conversation ID | Sets socket-local audio routing | +| `104` | JSON durable finalization request | Existing job and generation fencing applies | +| `105` | JSON speaker-sample request with optional `delivery_id` | Stable delivery is retained until opcode `202` on negotiated sockets | +| `107` | Header only | Requests bounded receiver drain while the acknowledgement socket remains open | +| `201` | JSON finalization result | Pusher to backend response | +| `202` | JSON `{kind, delivery_id}` | Pusher to backend acknowledgement after local worker invocation or a Redis done marker | + +Old pusher pods omit the capability response header. Listen then uses the legacy send-success +behavior for opcodes `102` and `105`, allowing pusher and listen to roll independently without +retaining frames forever against an old peer. + ### Server → Client (JSON over WS text frames) | Type | Format | Example | @@ -450,6 +502,9 @@ Staleness is fingerprint-driven: when late chunks change `audio_files` (pusher b | `PENDING_REQUEST_TIMEOUT` | 120s | `routers/listen/conversations.py` | Timeout before retrying a pending request | | `MAX_RETRIES_PER_REQUEST` | 3 | `routers/listen/conversations.py` | Max retries before keeping buffered | | `PUSHER_MAX_RECONNECT_ATTEMPTS` | 6 | `utils/listen_pusher_session.py` | Reconnect attempts before DEGRADED | +| `PUSHER_DELIVERY_ACK_TIMEOUT` | 5s | `utils/listen_pusher_session.py` | Retry interval for an unacknowledged stable delivery | +| `PUSHER_CLOSE_ACK_TIMEOUT` | 10s | `utils/listen_pusher_session.py` | Graceful-close acknowledgement drain | +| `PUSHER_SOCKET_CLOSE_TIMEOUT` | 2s | `utils/listen_pusher_session.py` | Socket-close deadline after the delivery drain | | `MAX_PENDING_REQUESTS` | 100 | `routers/listen/contracts.py` | Max buffered conversations per session | ## 8. WebSocket Task Supervision diff --git a/web/admin/lib/services/omi-api/omiApi.generated.ts b/web/admin/lib/services/omi-api/omiApi.generated.ts index 0f6694e748b..1529de10ece 100644 --- a/web/admin/lib/services/omi-api/omiApi.generated.ts +++ b/web/admin/lib/services/omi-api/omiApi.generated.ts @@ -1224,8 +1224,11 @@ export interface CreateTaskRequest { } export interface CreateTaskResponse { + ambiguous?: boolean | null; error?: string | null; + error_code?: string | null; external_task_id?: string | null; + retryable?: boolean | null; success: boolean; } diff --git a/web/app/src/lib/omiApi.generated.ts b/web/app/src/lib/omiApi.generated.ts index 0f6694e748b..1529de10ece 100644 --- a/web/app/src/lib/omiApi.generated.ts +++ b/web/app/src/lib/omiApi.generated.ts @@ -1224,8 +1224,11 @@ export interface CreateTaskRequest { } export interface CreateTaskResponse { + ambiguous?: boolean | null; error?: string | null; + error_code?: string | null; external_task_id?: string | null; + retryable?: boolean | null; success: boolean; } diff --git a/web/personas-open-source/src/lib/omiApi.generated.ts b/web/personas-open-source/src/lib/omiApi.generated.ts index 0f6694e748b..1529de10ece 100644 --- a/web/personas-open-source/src/lib/omiApi.generated.ts +++ b/web/personas-open-source/src/lib/omiApi.generated.ts @@ -1224,8 +1224,11 @@ export interface CreateTaskRequest { } export interface CreateTaskResponse { + ambiguous?: boolean | null; error?: string | null; + error_code?: string | null; external_task_id?: string | null; + retryable?: boolean | null; success: boolean; } From 350cdc891b7904874b757c50d65b25ab1f338f28 Mon Sep 17 00:00:00 2001 From: ZachL111 Date: Sat, 15 Aug 2026 05:40:38 -0700 Subject: [PATCH 2/3] Drop the superseded memory-convergence workflow contract and vector verification The legacy_memory_projection_convergence workflow entry belonged to the memory-stack half this PR dropped as superseded by the universal memory authority: its sources pulled database/memories.py into the tuple-result scan (flagging functions that live identically on main) and its test list referenced the deleted mcp search suite. The surviving vector_db delta was the same dropped surface's write-verification and returned a reported count where main's contract returns the payload length, failing main's batch-upsert tests; the file reverts to main wholesale. --- backend/database/vector_db.py | 19 +---- backend/testing/workflow_contracts.json | 103 +++++++++++++----------- 2 files changed, 57 insertions(+), 65 deletions(-) diff --git a/backend/database/vector_db.py b/backend/database/vector_db.py index 164495f98e6..6b11b29919b 100644 --- a/backend/database/vector_db.py +++ b/backend/database/vector_db.py @@ -389,14 +389,6 @@ class VectorCandidateQueryResult: rejected_count: int = 0 -def _reported_upsert_count(result: Any) -> int | None: - if isinstance(result, dict): - count = result.get('upserted_count') - else: - count = getattr(result, 'upserted_count', None) - return count if type(count) is int and count >= 0 else None - - def upsert_memory_vector( uid: str, memory_id: str, @@ -436,9 +428,6 @@ def upsert_memory_vector( } res = index.upsert(vectors=[data], namespace=MEMORIES_NAMESPACE) logger.info(f'upsert_memory_vector {memory_id} {res}') - if _reported_upsert_count(res) == 0: - logger.warning('upsert_memory_vector returned zero writes memory_id=%s', memory_id) - return None return vector @@ -494,8 +483,7 @@ def upsert_memory_vectors_batch(uid: str, items: List[Dict[str, Any]]) -> int: ) res = index.upsert(vectors=payload, namespace=MEMORIES_NAMESPACE) logger.info(f'upsert_memory_vectors_batch count={len(payload)} {res}') - reported_count = _reported_upsert_count(res) - return len(payload) if reported_count is None else reported_count + return len(payload) def find_similar_memories( @@ -661,18 +649,17 @@ def query_memory_vector_candidates( return VectorCandidateQueryResult(hits=hits, rejected_count=rejected_count) -def delete_memory_vector(uid: str, memory_id: str) -> bool: +def delete_memory_vector(uid: str, memory_id: str) -> None: """ Delete a memory vector from Pinecone. """ if index is None: logger.warning('Pinecone index not initialized, skipping memory vector delete') - return False + return vector_id = f'{uid}-{memory_id}' result = index.delete(ids=[vector_id], namespace=MEMORIES_NAMESPACE) logger.info(f'delete_memory_vector {vector_id} {result}') - return True def enqueue_projection_repair(uid: str, fact_id: str, reason: str, source_commit_id: str | None = None) -> List[str]: diff --git a/backend/testing/workflow_contracts.json b/backend/testing/workflow_contracts.json index 223d48c90e2..7af1c0ffcc1 100644 --- a/backend/testing/workflow_contracts.json +++ b/backend/testing/workflow_contracts.json @@ -130,7 +130,9 @@ "tests/unit/test_ws_c_backfill.py", "tests/unit/test_backfill_legacy_memories_cli.py" ], - "checks": ["no_large_tuple_results"], + "checks": [ + "no_large_tuple_results" + ], "invariants": [ "completed checkpoints recognize required-processing and quarantined pending-admission destinations without derived side effects", "partial checkpoints resume idempotently", @@ -144,12 +146,16 @@ { "id": "canonical_memory_fanout", "risk": "high", - "sources": ["backend/utils/memory/canonical_memory_adapter.py"], + "sources": [ + "backend/utils/memory/canonical_memory_adapter.py" + ], "tests": [ "tests/unit/test_ws_j_delete_privacy.py", "tests/unit/test_canonical_kg_promotion.py" ], - "checks": ["no_large_tuple_results"], + "checks": [ + "no_large_tuple_results" + ], "invariants": [ "domain writes use injected stores", "derived keyword/vector/KG fanout is either durable or explicitly recoverable" @@ -174,7 +180,9 @@ "tests/unit/test_ws_m_atom_keyword_index.py", "tests/unit/test_ws_n_graph_traversal.py" ], - "checks": ["no_large_tuple_results"], + "checks": [ + "no_large_tuple_results" + ], "invariants": [ "explicit submissions remain pending short-term until a content-bound processing receipt exists", "pending text is visible only to the first-party memory list and never to agent, developer, MCP, search, vector, or KG reads", @@ -250,7 +258,9 @@ "backend/scripts/sync_ledger_fence_cutover.py", ".github/workflows/sync_ledger_fence_cutover.yml" ], - "tests": ["tests/unit/test_sync_ledger_fence_cutover.py"], + "tests": [ + "tests/unit/test_sync_ledger_fence_cutover.py" + ], "checks": [], "invariants": [ "standby admission traffic reaches every sync surface before Cloud Tasks queues are paused", @@ -270,7 +280,9 @@ "tests/unit/test_vector_repair_outbox_worker.py", "tests/unit/test_vector_repair_outbox_infra.py" ], - "checks": ["no_large_tuple_results"], + "checks": [ + "no_large_tuple_results" + ], "invariants": [ "leased work is recoverable after worker death", "completed and dead-lettered work is never reclaimed" @@ -309,9 +321,15 @@ { "id": "projection_repair", "risk": "high", - "sources": ["backend/database/projection_repair.py"], - "tests": ["tests/unit/test_memory_ledger.py"], - "checks": ["no_large_tuple_results"], + "sources": [ + "backend/database/projection_repair.py" + ], + "tests": [ + "tests/unit/test_memory_ledger.py" + ], + "checks": [ + "no_large_tuple_results" + ], "invariants": [ "repair enqueue is idempotent", "repair processing uses injected stores and has retry/dead-letter accounting" @@ -339,7 +357,9 @@ "tests/unit/test_canonical_short_term_maintenance_cron.py", "tests/unit/test_memory_outbox_worker.py" ], - "checks": ["no_large_tuple_results"], + "checks": [ + "no_large_tuple_results" + ], "invariants": [ "every eligible short-term item receives exactly one terminal L2 route", "long-term admission and its graph assertion commit atomically against one exact short-term revision", @@ -372,9 +392,17 @@ { "id": "memory_ingestion_export_runner", "risk": "high", - "sources": ["backend/utils/memory_ingestion/export_runner.py", "backend/utils/memory_ingestion/pipeline.py"], - "tests": ["tests/unit/test_memory_ingestion_pipeline.py", "tests/unit/test_production_like_memory_model.py"], - "checks": ["no_large_tuple_results"], + "sources": [ + "backend/utils/memory_ingestion/export_runner.py", + "backend/utils/memory_ingestion/pipeline.py" + ], + "tests": [ + "tests/unit/test_memory_ingestion_pipeline.py", + "tests/unit/test_production_like_memory_model.py" + ], + "checks": [ + "no_large_tuple_results" + ], "invariants": [ "resume config mismatches fail closed", "failed shards remain visible in run summaries" @@ -410,10 +438,14 @@ "backend/utils/memory/**", "backend/utils/memory_ingestion/**" ], - "tests": ["testing/e2e/test_canonical_memory_pipeline.py"], - "checks": ["no_large_tuple_results"], + "tests": [ + "testing/e2e/test_canonical_memory_pipeline.py" + ], + "checks": [ + "no_large_tuple_results" + ], "invariants": [ - "capture→consolidate→promote→read with archive excluded from default reads", + "capture\u2192consolidate\u2192promote\u2192read with archive excluded from default reads", "canonical fail-closed without legacy bleed" ] }, @@ -449,7 +481,9 @@ "tests/unit/test_recording_sessions.py", "tests/unit/test_listen_finalization_cloud_tasks.py" ], - "checks": ["conversation-lifecycle-write-guard"], + "checks": [ + "conversation-lifecycle-write-guard" + ], "invariants": [ "only one finalizer claim succeeds for an in-progress conversation", "finalization outbox reconnects reuse one durable job and lease fences stale workers before any fanout", @@ -459,37 +493,6 @@ "discard and generic status-write escape hatches remain characterized until the lifecycle service replaces them" ] }, - { - "id": "legacy_memory_projection_convergence", - "risk": "high", - "sources": [ - "backend/database/memories.py", - "backend/database/memory_ledger.py", - "backend/database/projection_repair.py", - "backend/database/vector_db.py", - "backend/routers/memories.py", - "backend/utils/memory/memory_service.py", - "backend/utils/retrieval/tool_services/memories.py", - "backend/utils/x_connector.py" - ], - "tests": [ - "tests/unit/test_memory_service_parity.py", - "tests/unit/test_memories_batch_delete.py", - "tests/unit/test_memories_delete_batch_chunk.py", - "tests/unit/test_memory_ledger.py", - "tests/unit/test_memories_batch.py", - "tests/unit/test_mcp_search_memories.py", - "tests/unit/test_lock_bypass_fixes.py", - "tests/unit/test_tools_rest_memory_runtime_adapter.py", - "tests/unit/test_x_memory_extraction_retry.py" - ], - "checks": ["no_large_tuple_results"], - "invariants": [ - "legacy mutation guards, reads, and writes use one selected Firestore client", - "deletions enumerate every authoritative memory id before removing projections", - "search overfetches and filters stale, locked, rejected, invalid, or invalidated projections" - ] - }, { "id": "listen_pusher_pipeline", "risk": "high", @@ -676,7 +679,9 @@ "backend/charts/vad/**", "backend/scripts/validate_rendered_deployment_contract.py" ], - "tests": ["tests/unit/test_rendered_deployment_contract.py"], + "tests": [ + "tests/unit/test_rendered_deployment_contract.py" + ], "checks": [], "invariants": [ "every first-party GKE workload renders an explicit immutable image identity", From bde2b7ea6e5a0d8788eeb0ad9eea4611b59a1414 Mon Sep 17 00:00:00 2001 From: ZachL111 Date: Sat, 15 Aug 2026 07:40:17 -0700 Subject: [PATCH 3/3] Rerun against main with the desktop inventory drift fixed