Spaces:
Running
Running
File size: 5,644 Bytes
5c793cc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 | from app.modules.posts.domain.ports.embedding_provider import EmbeddingProvider
from app.modules.posts.domain.ports.post_sync_repository import PostSyncRepository
from app.modules.posts.domain.ports.post_sync_source import PostSyncSource
from app.modules.posts.domain.sync import PostEmbeddingRecord, PostSourceRecord, PostSyncResult
class SyncPostEmbeddingsUseCase:
CONSUMER = "post-embeddings-v1"
def __init__(
self,
source: PostSyncSource,
repository: PostSyncRepository,
embedding_provider: EmbeddingProvider,
embedding_model: str,
embedding_version: str,
embedding_dimension: int,
) -> None:
self._source = source
self._repository = repository
self._embedding_provider = embedding_provider
self._model = embedding_model
self._version = embedding_version
self._dimension = embedding_dimension
async def sync_snapshot(
self, batch_size: int = 100, page_limit: int | None = None, max_pages: int | None = None
) -> PostSyncResult:
result = PostSyncResult()
batch: list[PostSourceRecord] = []
async for post in self._source.iter_snapshot(page_limit, max_pages):
batch.append(post)
if len(batch) >= batch_size:
await self._upsert_batch(batch, result)
batch = []
await self._upsert_batch(batch, result)
return result
async def sync_incremental(
self, batch_size: int = 100, page_limit: int | None = None, max_pages: int | None = None
) -> PostSyncResult:
result = PostSyncResult()
after_id = await self._repository.get_checkpoint(self.CONSUMER)
pending: list[PostSourceRecord] = []
pending_last_event_id = after_id
async for change in self._source.iter_changes(after_id, page_limit, max_pages):
result.processed += 1
try:
if change.operation in {"delete", "archive", "deactivate"}:
if pending:
await self._upsert_batch(pending, result, count_processed=False)
await self._repository.save_checkpoint(
self.CONSUMER, pending_last_event_id
)
result.last_event_id = pending_last_event_id
pending = []
await self._repository.deactivate_embedding(
change.post_id, change.source_version
)
result.deactivated += 1
await self._repository.save_checkpoint(
self.CONSUMER, change.event_id
)
result.last_event_id = change.event_id
elif change.post is not None:
pending.append(change.post)
pending_last_event_id = change.event_id
if len(pending) >= batch_size:
await self._upsert_batch(pending, result, count_processed=False)
await self._repository.save_checkpoint(
self.CONSUMER, pending_last_event_id
)
result.last_event_id = pending_last_event_id
pending = []
else:
raise ValueError("un cambio upsert requiere el post completo")
except Exception as exc:
result.errors.append(f"event_id={change.event_id}: {exc}")
break
if pending and not result.errors:
await self._upsert_batch(pending, result, count_processed=False)
await self._repository.save_checkpoint(
self.CONSUMER, pending_last_event_id
)
result.last_event_id = pending_last_event_id
return result
async def _upsert_batch(
self,
batch: list[PostSourceRecord],
result: PostSyncResult,
*,
count_processed: bool = True,
) -> None:
if not batch:
return
if count_processed:
result.processed += len(batch)
hashes = await self._repository.fetch_content_hashes(item.id for item in batch)
missing_versions = [item.id for item in batch if item.source_version is None]
if missing_versions:
raise ValueError(
"source_version es obligatorio para posts: "
+ ",".join(missing_versions[:20])
)
expected = {
item.id: self._versioned_hash(item.content_hash, item.source_version)
for item in batch
}
changed = [item for item in batch if hashes.get(item.id) != expected[item.id]]
result.skipped += len(batch) - len(changed)
if not changed:
return
vectors = self._embedding_provider.embed_batch([item.document for item in changed])
if len(vectors) != len(changed):
raise ValueError("el proveedor de embeddings devolvio un lote incompleto")
records = [
PostEmbeddingRecord(item, vector, expected[item.id])
for item, vector in zip(changed, vectors)
]
await self._repository.upsert_embeddings(records)
result.upserted += len(records)
def _versioned_hash(self, source_hash: str, source_version: int | None) -> str:
from hashlib import sha256
raw = (
f"{source_hash}:{self._model}:{self._version}:"
f"{self._dimension}:{source_version}"
)
return sha256(raw.encode("utf-8")).hexdigest()
|