Skip to content
12 changes: 12 additions & 0 deletions deployments/entity-linkage/src/stitch/entity_linkage/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,18 @@ async def iter_oil_gas_fields(
):
yield self._to_candidate(item)

async def get_oil_gas_fields_total(self) -> int | None:
"""Total resource count, for use as a linkage-progress denominator.

Fetches a single-item page purely to read ``total_count``; the streaming
iterator used by the pass itself discards that field. Returns ``None`` if
the payload omits an integer count, so progress can still report a
numerator without a denominator.
"""
payload = await self._client.list_oil_gas_fields_page(page=1, page_size=1)
total = payload.get("total_count")
return total if isinstance(total, int) else None

async def list_merge_candidates(self) -> list[dict[str, Any]]:
return await self._client.list_merge_candidates()

Expand Down
18 changes: 18 additions & 0 deletions deployments/entity-linkage/src/stitch/entity_linkage/entities.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,24 @@ class BulkLinkResponse(BaseModel):
resources_failed: int = 0


class LinkProgress(BaseModel):
"""In-flight progress of a running linkage pass.

Written onto the live job record as the pass streams resources, so a poller
can see how far along a multi-hour run is rather than only "running". Field
names mirror :class:`BulkLinkResponse` so the running view and the final
result read consistently. ``total_resources`` is ``None`` when the
denominator could not be fetched; percent is derived by the caller.
"""

resources_scanned: int
total_resources: int | None
merge_candidates_created: int
merge_candidates_skipped: int
resources_failed: int
updated_at: datetime


class PaginationParams(BaseModel):
page: int = Field(1, ge=1)
page_size: int = Field(50, ge=1, le=200)
Expand Down
8 changes: 5 additions & 3 deletions deployments/entity-linkage/src/stitch/entity_linkage/jobs.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@

logger = logging.getLogger("stitch.entity_linkage")

RunThunk = Callable[[], Awaitable[BaseModel]]
RunThunk = Callable[["JobRecord"], Awaitable[BaseModel]]
Comment thread
AlexAxthelm marked this conversation as resolved.


class JobState(str, Enum):
Expand All @@ -39,6 +39,7 @@ class JobRecord(BaseModel):
params: SerializeAsAny[BaseModel]
started_at: datetime
finished_at: datetime | None = None
progress: SerializeAsAny[BaseModel] | None = None
result: SerializeAsAny[BaseModel] | None = None
error: str | None = None

Expand All @@ -53,7 +54,8 @@ class JobManager:
"""Single-job, in-memory run manager.

State is lost on restart and concurrent runs are rejected. The run body is
supplied per start as a zero-arg coroutine, so this manager is generic.
supplied per start as a coroutine that receives the live ``JobRecord``, so it
can publish progress onto the record while it runs; this manager stays generic.
"""

def __init__(self) -> None:
Expand Down Expand Up @@ -81,7 +83,7 @@ async def start(self, params: BaseModel, run: RunThunk) -> JobRecord:

async def _run(self, record: JobRecord, run: RunThunk) -> None:
try:
record.result = await run()
record.result = await run(record)
record.state = JobState.succeeded
except Exception as exc:
logger.exception("Linkage run %s failed", record.job_id)
Expand Down
50 changes: 49 additions & 1 deletion deployments/entity-linkage/src/stitch/entity_linkage/matching.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,13 +21,15 @@
from __future__ import annotations

import logging
from collections.abc import Sequence
from collections.abc import Callable, Sequence
from datetime import UTC, datetime

import httpx

from stitch.entity_linkage.client import StitchApiClient
from stitch.entity_linkage.entities import (
BulkLinkResponse,
LinkProgress,
ResourceLinkResult,
normalize_country,
normalize_name,
Expand All @@ -36,6 +38,11 @@

logger = logging.getLogger(__name__)

# How often the bulk pass publishes a progress snapshot, in resources scanned.
# Small enough that a 2s poller sees frequent movement, large enough to avoid
# building a progress model on every one of hundreds of thousands of resources.
PROGRESS_UPDATE_EVERY = 100

# A 4xx from create-merge-candidate is an expected, non-fatal outcome during a
# run: the API rejects a duplicate fingerprint (a candidate already exists) or a
# resource that has already been merged. We skip those rather than aborting.
Expand Down Expand Up @@ -171,26 +178,64 @@ async def link_all(
apply_merges: bool,
page_size: int,
initiated_by: str,
on_progress: Callable[[LinkProgress], None] | None = None,
) -> BulkLinkResponse:
"""Run the bounded matcher over every resource, streaming ids page by page.

Groups are de-duplicated by fingerprint across the run, so each block is
submitted at most once even though every member rediscovers it. Members of an
already-formed block are skipped without re-searching.

``on_progress``, when supplied, is called with a :class:`LinkProgress`
snapshot periodically (every ``PROGRESS_UPDATE_EVERY`` resources) and once
more at the end, so a poller can track a long run's advance.
"""
# Only needed when we will actually POST; skip the (currently unpaginated)
# candidate-list fetch entirely on a dry run.
known_existing = await _existing_fingerprints(client) if apply_merges else None

# Denominator for progress; only worth an extra request when a progress
# consumer is listening. None if unavailable -- a failure here must not abort
# the pass, so fall back to an unknown total.
total_resources: int | None = None
if on_progress is not None:
try:
total_resources = await client.get_oil_gas_fields_total()
except (StitchAPIError, httpx.HTTPError, OSError) as exc:
logger.warning("Could not fetch resource total for progress: %s", exc)
total_resources = None

groups_by_fingerprint: dict[str, list[int]] = {}
processed_ids: set[int] = set()
resources_scanned = 0
created = 0
skipped = 0
failed = 0

def emit_progress() -> None:
if on_progress is None:
return
on_progress(
LinkProgress(
resources_scanned=resources_scanned,
total_resources=total_resources,
merge_candidates_created=created,
merge_candidates_skipped=skipped,
resources_failed=failed,
updated_at=datetime.now(UTC),
)
)

# Publish a 0/total snapshot up front so a poller sees a denominator (and any
# progress at all) before the 100th resource -- and, for a run shorter than
# one throttle window, at all, since the state flips to succeeded right after
# the final snapshot with no yield in between.
emit_progress()

async for candidate in client.iter_oil_gas_fields(page_size=page_size):
resources_scanned += 1
if resources_scanned % PROGRESS_UPDATE_EVERY == 0:
emit_progress()
if candidate.id in processed_ids:
continue

Expand Down Expand Up @@ -232,6 +277,9 @@ async def link_all(
elif was_skipped:
skipped += 1

# Final snapshot so the last poll before completion reflects the exact totals.
emit_progress()

return BulkLinkResponse(
initiated_by=initiated_by,
apply_merges=apply_merges,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -83,13 +83,14 @@ async def start_link_all(
"""
initiated_by = user_label(auth_context.user)

async def run() -> BulkLinkResponse:
async def run(record: JobRecord) -> BulkLinkResponse:
async with StitchApiClient() as client:
return await matching.link_all(
client,
apply_merges=request.apply_merges,
page_size=request.page_size,
initiated_by=initiated_by,
on_progress=lambda progress: setattr(record, "progress", progress),
)

try:
Expand Down
31 changes: 28 additions & 3 deletions deployments/entity-linkage/tests/test_jobs.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

from stitch.entity_linkage.jobs import (
JobAlreadyRunningError,
JobRecord,
JobState,
get_job_manager,
reset_manager,
Expand All @@ -21,6 +22,10 @@ class _Result(BaseModel):
doubled: int


class _Progress(BaseModel):
scanned: int


@pytest.fixture(autouse=True)
def _reset():
reset_manager()
Expand All @@ -42,7 +47,7 @@ def test_run_thunk_success_records_result() -> None:
async def scenario() -> None:
mgr = get_job_manager()

async def run() -> _Result:
async def run(_record: JobRecord) -> _Result:
return _Result(doubled=6)

record = await mgr.start(_Params(n=3), run)
Expand All @@ -63,7 +68,7 @@ def test_run_thunk_failure_records_error() -> None:
async def scenario() -> None:
mgr = get_job_manager()

async def run() -> _Result:
async def run(_record: JobRecord) -> _Result:
raise RuntimeError("kaboom")

await mgr.start(_Params(), run)
Expand All @@ -82,7 +87,7 @@ def test_manager_rejects_concurrent_start() -> None:
async def scenario() -> None:
mgr = get_job_manager()

async def slow() -> _Result:
async def slow(_record: JobRecord) -> _Result:
await asyncio.sleep(0.5)
return _Result(doubled=0)

Expand All @@ -91,3 +96,23 @@ async def slow() -> _Result:
await mgr.start(_Params(), slow)

asyncio.run(scenario())


def test_run_thunk_can_write_progress_onto_record() -> None:
async def scenario() -> None:
mgr = get_job_manager()

async def run(record: JobRecord) -> _Result:
# The run body writes progress onto the live record; a poller reading
# mgr.current() must see it while the run is still in flight.
record.progress = _Progress(scanned=42)
return _Result(doubled=0)

await mgr.start(_Params(), run)
for _ in range(200):
if mgr.current().state != JobState.running:
break
await asyncio.sleep(0.01)
assert mgr.current().progress.model_dump() == {"scanned": 42}

asyncio.run(scenario())
15 changes: 14 additions & 1 deletion deployments/entity-linkage/tests/test_link_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,9 @@ async def get_oil_gas_field_detail(self, resource_id: int) -> FieldDetailCandida
raise self.detail_error
return self.details_by_id[resource_id]

async def get_oil_gas_fields_total(self) -> int | None:
return len(self.items)

async def iter_oil_gas_fields(
self,
*,
Expand Down Expand Up @@ -282,6 +285,14 @@ def test_link_all_launches_job_and_status_succeeds(test_client, install_client)
assert result["merge_candidates_skipped"] == 0
assert fake.create_calls == [[1, 2]]

# The run publishes a final progress snapshot onto the record.
progress = final["progress"]
assert progress is not None
assert progress["resources_scanned"] == 3
assert progress["total_resources"] == 3
assert progress["merge_candidates_created"] == 1
assert progress["updated_at"] is not None


def test_link_all_records_downstream_failure_in_status(
test_client, install_client
Expand Down Expand Up @@ -311,7 +322,9 @@ def test_link_all_rejects_concurrent_run_with_409(
) -> None:
install_client()

async def slow_link_all(client, *, apply_merges, page_size, initiated_by):
async def slow_link_all(
client, *, apply_merges, page_size, initiated_by, on_progress=None
):
import asyncio

await asyncio.sleep(0.5)
Expand Down
Loading
Loading