|
3 | 3 | import time |
4 | 4 | from collections.abc import Sequence |
5 | 5 | from dataclasses import dataclass, field |
| 6 | +from datetime import datetime, timezone |
6 | 7 | from typing import TYPE_CHECKING, Any, Literal |
7 | 8 |
|
8 | 9 | from zep_ingest._validation import require_int_range, require_nonnegative_number |
|
41 | 42 | ] |
42 | 43 |
|
43 | 44 |
|
| 45 | +def _parse_created_at(value: str | None) -> datetime | None: |
| 46 | + """Parse an episode ``created_at`` into a comparable UTC instant.""" |
| 47 | + if not value: |
| 48 | + return None |
| 49 | + try: |
| 50 | + parsed = datetime.fromisoformat(value.replace("Z", "+00:00")) |
| 51 | + except ValueError: |
| 52 | + return None |
| 53 | + if parsed.tzinfo is None: |
| 54 | + return parsed.replace(tzinfo=timezone.utc) |
| 55 | + return parsed |
| 56 | + |
| 57 | + |
| 58 | +def _is_later_or_equal_created_at(new_at: str | None, prior_at: str | None) -> bool: |
| 59 | + """Whether ``new_at`` is the later document tail. |
| 60 | +
|
| 61 | + Missing or unparseable timestamps fall back to last-submitted-wins. Offsets |
| 62 | + are compared as instants, not lexicographic RFC3339 strings. |
| 63 | + """ |
| 64 | + new_dt = _parse_created_at(new_at) |
| 65 | + prior_dt = _parse_created_at(prior_at) |
| 66 | + if new_dt is None or prior_dt is None: |
| 67 | + return True |
| 68 | + return new_dt >= prior_dt |
| 69 | + |
| 70 | + |
44 | 71 | def _normalize_task_status(status: str | None) -> str: |
45 | 72 | status = status.lower() if status is not None else None |
46 | 73 | if status is None or status in {"created", "draft", "pending", "queued"}: |
@@ -178,11 +205,7 @@ def record_sequential_episode(self, episode: Episode, uuid: str) -> None: |
178 | 205 | self._uses_document_grouping = True |
179 | 206 | prior_at = self._document_poll_created_at.get(episode.document_id) |
180 | 207 | new_at = episode.created_at |
181 | | - if prior_at is None or new_at is None: |
182 | | - # Missing timestamps: fall back to submission order (last wins). |
183 | | - self._document_poll_uuids[episode.document_id] = uuid |
184 | | - self._document_poll_created_at[episode.document_id] = new_at |
185 | | - elif new_at >= prior_at: |
| 208 | + if _is_later_or_equal_created_at(new_at, prior_at): |
186 | 209 | self._document_poll_uuids[episode.document_id] = uuid |
187 | 210 | self._document_poll_created_at[episode.document_id] = new_at |
188 | 211 | else: |
@@ -362,10 +385,7 @@ def combine(self, *others: "IngestResult") -> "IngestResult": |
362 | 385 | for document_id, uuid in part._document_poll_uuids.items(): |
363 | 386 | created_at = part._document_poll_created_at.get(document_id) |
364 | 387 | prior_at = combined._document_poll_created_at.get(document_id) |
365 | | - if prior_at is None or created_at is None: |
366 | | - combined._document_poll_uuids[document_id] = uuid |
367 | | - combined._document_poll_created_at[document_id] = created_at |
368 | | - elif (created_at or "") >= (prior_at or ""): |
| 388 | + if _is_later_or_equal_created_at(created_at, prior_at): |
369 | 389 | combined._document_poll_uuids[document_id] = uuid |
370 | 390 | combined._document_poll_created_at[document_id] = created_at |
371 | 391 | if part._plain_episode_tail is not None: |
|
0 commit comments