Source code for rowguard.results.async_stream_result

from __future__ import annotations

from collections.abc import AsyncIterator, Sequence
from contextlib import suppress
from time import perf_counter_ns
from typing import Any, Generic, TypeVar

from pydantic import BaseModel

from rowguard.diagnostics import Diagnostic
from rowguard.errors import QueryExecutionError, RowGuardError
from rowguard.execution.async_ import aclose_result
from rowguard.execution.context import AsyncExecutionContext
from rowguard.execution.guards import require_session_for_entity_plan
from rowguard.execution.observer import StreamObserver
from rowguard.execution.processor import aprocess_row, build_rejection_context
from rowguard.execution.state import MutableStatistics
from rowguard.execution.thresholds import check_rejection_thresholds
from rowguard.planning.config import StreamingConfig
from rowguard.planning.execution_plan import ExecutionPlan
from rowguard.results.quarantine import QuarantineReceipt
from rowguard.results.rejected_row import RejectedRow
from rowguard.statistics import QueryStatistics

T = TypeVar("T", bound=BaseModel)


[docs] class AsyncStreamResult(Generic[T]): """Async incremental validated-row iterator with context-managed DB cleanup. Accepted models are yielded and never retained. Prefer ``async with rowguard.astream(...) as stream:`` or ``async for model in stream``. """ def __init__( self, *, plan: ExecutionPlan[T], context: AsyncExecutionContext, streaming: StreamingConfig | None = None, observers: Sequence[StreamObserver] = (), ) -> None: self._plan = plan self._context = context self._streaming = streaming or StreamingConfig() self._observers: tuple[StreamObserver, ...] = tuple(observers) self._statistics = MutableStatistics() self._rejected: list[RejectedRow] = [] self._quarantine_receipts: list[QuarantineReceipt] = [] self._diagnostics: list[Diagnostic] = list(plan.diagnostics) self._db_result: Any | None = None self._row_aiter: AsyncIterator[Any] | None = None self._row_iter: Any | None = None self._index = 0 self._started = False self._closed = False self._completed = False self._started_ns = 0 self._primary_error: BaseException | None = None @property def closed(self) -> bool: return self._closed @property def statistics(self) -> QueryStatistics: return self._statistics.snapshot() @property def rejected(self) -> tuple[RejectedRow, ...]: return tuple(self._rejected) @property def quarantine_receipts(self) -> tuple[QuarantineReceipt, ...]: return tuple(self._quarantine_receipts) @property def diagnostics(self) -> tuple[Diagnostic, ...]: return tuple(self._diagnostics) @property def statement(self) -> Any: return self._plan.statement @property def has_rejections(self) -> bool: return self._statistics.rows_rejected > 0 @property def is_clean(self) -> bool: return not self.has_rejections @property def rejected_count(self) -> int: return len(self._rejected) @property def execution_time(self) -> float: return self._statistics.execution_time_ns / 1_000_000_000 def __aiter__(self) -> AsyncIterator[T]: return self._iterate() async def _iterate(self) -> AsyncIterator[T]: await self._ensure_started() try: while True: yield await self._next_model() except StopAsyncIteration: return finally: await self.close() async def __aenter__(self) -> AsyncStreamResult[T]: await self._ensure_started() return self async def __aexit__( self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: object, ) -> None: if exc is not None and self._primary_error is None: self._primary_error = exc self._notify_failed(exc) await self.close() async def __anext__(self) -> T: await self._ensure_started() return await self._next_model()
[docs] async def close(self) -> None: if self._closed: return self._closed = True self._stamp_execution_time() close_error: BaseException | None = None if self._db_result is not None: try: await aclose_result(self._db_result) except Exception as error: close_error = error self._db_result = None self._row_aiter = None self._row_iter = None aclose = getattr(self._plan.rejection_policy, "aclose", None) if callable(aclose): try: await aclose() except Exception as error: if close_error is None: close_error = error else: sync_close = getattr(self._plan.rejection_policy, "close", None) if callable(sync_close): try: sync_close() except Exception as error: if close_error is None: close_error = error self._notify_closed() if close_error is not None and self._primary_error is None: self._primary_error = close_error raise close_error
async def _next_model(self) -> T: if self._closed: raise StopAsyncIteration assert self._row_aiter is not None or self._row_iter is not None while True: try: if self._row_aiter is not None: row = await self._row_aiter.__anext__() else: assert self._row_iter is not None try: row = next(self._row_iter) except StopIteration as stop: raise StopAsyncIteration from stop except StopAsyncIteration: await self._finish_complete() raise index = self._index self._index += 1 try: rejection_context = build_rejection_context( plan=self._plan, rows_read=self._statistics.rows_read, rows_accepted=self._statistics.rows_accepted, rows_rejected=self._statistics.rows_rejected, ) processed = await aprocess_row( row=row, index=index, plan=self._plan, context=rejection_context, ) except RowGuardError as error: self._primary_error = error self._notify_failed(error) await self.close() raise except Exception as error: wrapped = QueryExecutionError(f"Query execution failed: {error}") wrapped.__cause__ = error self._primary_error = wrapped self._notify_failed(wrapped) await self.close() raise wrapped from error self._statistics.record_processed(processed) if processed.model is not None: self._notify_accepted(index, processed.model) return processed.model if processed.quarantine_receipt is not None: self._quarantine_receipts.append(processed.quarantine_receipt) if processed.rejected is not None: self._notify_rejected(processed.rejected) if processed.retain_rejection: self._rejected.append(processed.rejected) if processed.raise_error is not None: self._primary_error = processed.raise_error self._notify_failed(processed.raise_error) await self.close() raise processed.raise_error try: check_rejection_thresholds( statistics=self._statistics, rejection_plan=self._plan.rejection_plan, last_rejection=processed.rejected, ) except RowGuardError as error: self._primary_error = error self._notify_failed(error) await self.close() raise if not processed.continue_processing: await self._finish_complete() raise StopAsyncIteration async def _finish_complete(self) -> None: if self._completed: await self.close() return self._completed = True self._stamp_execution_time() self._notify_complete() await self.close() def _stamp_execution_time(self) -> None: if self._started_ns: self._statistics.execution_time_ns = perf_counter_ns() - self._started_ns async def _ensure_started(self) -> None: if self._closed: raise QueryExecutionError("AsyncStreamResult is closed and cannot be reused") if self._started: return self._started = True self._started_ns = perf_counter_ns() self._notify_start() try: self._db_result = await self._stream_statement() if hasattr(self._db_result, "__aiter__"): self._row_aiter = self._db_result.__aiter__() self._row_iter = None else: # CursorResult / sync-iterable fallback (parity with AsyncExecutionEngine). self._row_aiter = None self._row_iter = iter(self._db_result) except RowGuardError as error: self._primary_error = error self._notify_failed(error) await self.close() raise except Exception as error: wrapped = QueryExecutionError(f"Query execution failed: {error}") wrapped.__cause__ = error self._primary_error = wrapped self._notify_failed(wrapped) await self.close() raise wrapped from error async def _stream_statement(self) -> Any: require_session_for_entity_plan(self._plan, session=self._context.session) statement = self._plan.statement options: dict[str, Any] = {} if self._streaming.stream_results: options["stream_results"] = True if self._streaming.yield_per is not None: options["yield_per"] = self._streaming.yield_per if options: statement = statement.execution_options(**options) params = dict(self._plan.parameters) if self._plan.parameters else {} if self._context.session is not None: target = self._context.session else: target = self._context.connection if target is None: raise QueryExecutionError("No session or connection available for execution") stream = getattr(target, "stream", None) if callable(stream): if params: return await stream(statement, params) return await stream(statement) # Fallback for unusual async handles that only expose execute(). if params: return await target.execute(statement, params) return await target.execute(statement) def _notify_start(self) -> None: for observer in self._observers: try: observer.on_stream_start(execution_id=self._plan.execution_id) except Exception as error: self._handle_observer_error(error) def _notify_accepted(self, index: int, model: T) -> None: for observer in self._observers: try: observer.on_row_accepted(index=index, model=model) except Exception as error: self._handle_observer_error(error) def _notify_rejected(self, rejected: RejectedRow) -> None: for observer in self._observers: try: observer.on_row_rejected(rejected=rejected) except Exception as error: self._handle_observer_error(error) def _notify_complete(self) -> None: stats = self._statistics.snapshot() for observer in self._observers: try: observer.on_stream_complete(statistics=stats) except Exception as error: self._handle_observer_error(error) def _notify_failed(self, error: BaseException) -> None: for observer in self._observers: try: observer.on_stream_failed(error=error) except Exception as observer_error: self._handle_observer_error(observer_error) def _notify_closed(self) -> None: for observer in self._observers: try: observer.on_stream_closed() except Exception as error: if self._primary_error is None: self._handle_observer_error(error) else: with suppress(Exception): self._handle_observer_error(error) def _handle_observer_error(self, error: Exception) -> None: self._diagnostics.append( Diagnostic( code="streaming.observer_error", severity="warning", execution_id=self._plan.execution_id, metadata={"error": str(error), "error_type": type(error).__name__}, ) )