additional config, pollers, serve
This commit is contained in:
parent
bdcc1caa49
commit
75c3896588
15 changed files with 1286 additions and 7 deletions
|
|
@ -4,10 +4,14 @@ from haiku.rag.config.loader import (
|
||||||
load_yaml_config,
|
load_yaml_config,
|
||||||
)
|
)
|
||||||
from haiku.rag.config.models import (
|
from haiku.rag.config.models import (
|
||||||
|
APIConfig,
|
||||||
AppConfig,
|
AppConfig,
|
||||||
|
CircuitBreakerConfig,
|
||||||
ConversionOptions,
|
ConversionOptions,
|
||||||
EmbeddingModelConfig,
|
EmbeddingModelConfig,
|
||||||
EmbeddingsConfig,
|
EmbeddingsConfig,
|
||||||
|
FSSourceConfig,
|
||||||
|
HTTPSourceConfig,
|
||||||
IngesterConfig,
|
IngesterConfig,
|
||||||
LanceDBConfig,
|
LanceDBConfig,
|
||||||
ModelConfig,
|
ModelConfig,
|
||||||
|
|
@ -19,16 +23,24 @@ from haiku.rag.config.models import (
|
||||||
QAConfig,
|
QAConfig,
|
||||||
QueueConfig,
|
QueueConfig,
|
||||||
RerankingConfig,
|
RerankingConfig,
|
||||||
|
RetryPolicyConfig,
|
||||||
S3MonitorEntry,
|
S3MonitorEntry,
|
||||||
|
S3SourceConfig,
|
||||||
|
SourceConfig,
|
||||||
StorageConfig,
|
StorageConfig,
|
||||||
|
WorkerConfig,
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"Config",
|
"Config",
|
||||||
|
"APIConfig",
|
||||||
"AppConfig",
|
"AppConfig",
|
||||||
|
"CircuitBreakerConfig",
|
||||||
"ConversionOptions",
|
"ConversionOptions",
|
||||||
"EmbeddingModelConfig",
|
"EmbeddingModelConfig",
|
||||||
"EmbeddingsConfig",
|
"EmbeddingsConfig",
|
||||||
|
"FSSourceConfig",
|
||||||
|
"HTTPSourceConfig",
|
||||||
"IngesterConfig",
|
"IngesterConfig",
|
||||||
"LanceDBConfig",
|
"LanceDBConfig",
|
||||||
"ModelConfig",
|
"ModelConfig",
|
||||||
|
|
@ -40,8 +52,12 @@ __all__ = [
|
||||||
"QAConfig",
|
"QAConfig",
|
||||||
"QueueConfig",
|
"QueueConfig",
|
||||||
"RerankingConfig",
|
"RerankingConfig",
|
||||||
|
"RetryPolicyConfig",
|
||||||
"S3MonitorEntry",
|
"S3MonitorEntry",
|
||||||
|
"S3SourceConfig",
|
||||||
|
"SourceConfig",
|
||||||
"StorageConfig",
|
"StorageConfig",
|
||||||
|
"WorkerConfig",
|
||||||
"find_config_file",
|
"find_config_file",
|
||||||
"generate_default_config",
|
"generate_default_config",
|
||||||
"get_config",
|
"get_config",
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,5 @@
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Literal
|
from typing import Annotated, Literal
|
||||||
|
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
|
@ -258,10 +258,101 @@ class QueueConfig(BaseModel):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class IngesterConfig(BaseModel):
|
class RetryPolicyConfig(BaseModel):
|
||||||
"""Production ingester settings. Expanded across chunks 4-7."""
|
"""Per-job retry policy. Per-source override is allowed under
|
||||||
|
SourceConfig.retry so a flaky source doesn't drag the rest of the queue."""
|
||||||
|
|
||||||
|
max_attempts: int = 5
|
||||||
|
base_delay_s: float = 2.0
|
||||||
|
max_delay_s: float = 300.0
|
||||||
|
jitter: float = Field(default=0.25, ge=0.0, le=1.0)
|
||||||
|
|
||||||
|
|
||||||
|
class CircuitBreakerConfig(BaseModel):
|
||||||
|
"""Per-source breaker over discover() failures. Stops the ingester from
|
||||||
|
hammering a source that's persistently failing."""
|
||||||
|
|
||||||
|
failure_threshold: int = Field(
|
||||||
|
default=5, description="Consecutive failures before the breaker opens."
|
||||||
|
)
|
||||||
|
cooldown_s: float = Field(
|
||||||
|
default=600.0,
|
||||||
|
description="How long the breaker stays open before allowing a probe.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class WorkerConfig(BaseModel):
|
||||||
|
worker_count: int = 4
|
||||||
|
max_concurrent: int = 4
|
||||||
|
poll_idle_interval_s: float = 1.0
|
||||||
|
claim_timeout_s: int = 1800
|
||||||
|
reaper_interval_s: int = 60
|
||||||
|
retry: RetryPolicyConfig = Field(default_factory=RetryPolicyConfig)
|
||||||
|
|
||||||
|
|
||||||
|
class APIConfig(BaseModel):
|
||||||
|
"""HTTP control plane settings for the ingester."""
|
||||||
|
|
||||||
|
enabled: bool = True
|
||||||
|
host: str = "127.0.0.1"
|
||||||
|
port: int = 8765
|
||||||
|
auth_token: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class _SourceBase(BaseModel):
|
||||||
|
"""Fields common to every source. `id` is optional; if omitted the source
|
||||||
|
derives a deterministic id from its target (root path / bucket+prefix /
|
||||||
|
user-supplied tag)."""
|
||||||
|
|
||||||
|
id: str | None = None
|
||||||
|
delete_orphans: bool = True
|
||||||
|
poll_interval_s: float = Field(
|
||||||
|
default=300.0,
|
||||||
|
description="How often discover() runs. FS additionally uses watchfiles "
|
||||||
|
"for push events between sweeps.",
|
||||||
|
)
|
||||||
|
retry: RetryPolicyConfig | None = Field(
|
||||||
|
default=None,
|
||||||
|
description="Override the worker's default retry policy for jobs from "
|
||||||
|
"this source. None = inherit from WorkerConfig.retry.",
|
||||||
|
)
|
||||||
|
circuit_breaker: CircuitBreakerConfig = Field(default_factory=CircuitBreakerConfig)
|
||||||
|
|
||||||
|
|
||||||
|
class FSSourceConfig(_SourceBase):
|
||||||
|
type: Literal["fs"]
|
||||||
|
root: Path
|
||||||
|
ignore_patterns: list[str] = []
|
||||||
|
include_patterns: list[str] = []
|
||||||
|
|
||||||
|
|
||||||
|
class HTTPSourceConfig(_SourceBase):
|
||||||
|
type: Literal["http"]
|
||||||
|
urls: list[str] = []
|
||||||
|
headers: dict[str, str] = Field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
class S3SourceConfig(_SourceBase):
|
||||||
|
type: Literal["s3"]
|
||||||
|
uri: str
|
||||||
|
storage_options: dict[str, str] = Field(default_factory=dict)
|
||||||
|
ignore_patterns: list[str] = []
|
||||||
|
include_patterns: list[str] = []
|
||||||
|
|
||||||
|
|
||||||
|
SourceConfig = Annotated[
|
||||||
|
FSSourceConfig | HTTPSourceConfig | S3SourceConfig,
|
||||||
|
Field(discriminator="type"),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class IngesterConfig(BaseModel):
|
||||||
|
"""Production ingester settings."""
|
||||||
|
|
||||||
|
sources: list[SourceConfig] = []
|
||||||
queue: QueueConfig = Field(default_factory=QueueConfig)
|
queue: QueueConfig = Field(default_factory=QueueConfig)
|
||||||
|
workers: WorkerConfig = Field(default_factory=WorkerConfig)
|
||||||
|
api: APIConfig = Field(default_factory=APIConfig)
|
||||||
|
|
||||||
|
|
||||||
class AppConfig(BaseModel):
|
class AppConfig(BaseModel):
|
||||||
|
|
|
||||||
110
haiku_rag_slim/haiku/rag/ingester/app.py
Normal file
110
haiku_rag_slim/haiku/rag/ingester/app.py
Normal file
|
|
@ -0,0 +1,110 @@
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import signal
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from haiku.rag.config import AppConfig
|
||||||
|
from haiku.rag.ingester.pollers.manager import PollerManager
|
||||||
|
from haiku.rag.ingester.queue.migrations import open_queue
|
||||||
|
from haiku.rag.ingester.queue.repository import JobRepo, SyncStateRepo
|
||||||
|
from haiku.rag.ingester.workers.pool import WorkerPool
|
||||||
|
from haiku.rag.ingester.workers.retry import RetryPolicy
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import aiosqlite
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class IngesterApp:
|
||||||
|
"""Top-level lifecycle for the production ingester.
|
||||||
|
|
||||||
|
Owns: SQLite queue connection, JobRepo/SyncStateRepo, PollerManager,
|
||||||
|
WorkerPool, and a HaikuRAG client for the worker pool to ingest through.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, *, config: AppConfig, db_path: Path):
|
||||||
|
self._config = config
|
||||||
|
self._db_path = db_path
|
||||||
|
self._queue_conn: aiosqlite.Connection | None = None
|
||||||
|
self._jobs: JobRepo | None = None
|
||||||
|
self._sync: SyncStateRepo | None = None
|
||||||
|
self._pool: WorkerPool | None = None
|
||||||
|
self._pollers: PollerManager | None = None
|
||||||
|
self._client = None
|
||||||
|
|
||||||
|
async def serve(self, *, api: bool = True) -> None:
|
||||||
|
"""Run pollers + workers (and the HTTP API when enabled) until a
|
||||||
|
SIGINT/SIGTERM is received. Drains in-flight work on shutdown."""
|
||||||
|
from haiku.rag.client import HaikuRAG
|
||||||
|
from haiku.rag.converters import get_converter
|
||||||
|
|
||||||
|
ingester_cfg = self._config.ingester
|
||||||
|
self._queue_conn = await open_queue(ingester_cfg.queue.path)
|
||||||
|
self._jobs = JobRepo(self._queue_conn)
|
||||||
|
self._sync = SyncStateRepo(self._queue_conn)
|
||||||
|
|
||||||
|
supported_extensions = get_converter(self._config).supported_extensions
|
||||||
|
retry = RetryPolicy(
|
||||||
|
max_attempts=ingester_cfg.workers.retry.max_attempts,
|
||||||
|
base_delay_s=ingester_cfg.workers.retry.base_delay_s,
|
||||||
|
max_delay_s=ingester_cfg.workers.retry.max_delay_s,
|
||||||
|
jitter=ingester_cfg.workers.retry.jitter,
|
||||||
|
)
|
||||||
|
|
||||||
|
async with HaikuRAG(self._db_path, config=self._config) as client:
|
||||||
|
self._client = client
|
||||||
|
self._pool = WorkerPool(
|
||||||
|
client=client,
|
||||||
|
job_repo=self._jobs,
|
||||||
|
sync_repo=self._sync,
|
||||||
|
worker_count=ingester_cfg.workers.worker_count,
|
||||||
|
max_concurrent=ingester_cfg.workers.max_concurrent,
|
||||||
|
retry_policy=retry,
|
||||||
|
poll_idle_interval_s=ingester_cfg.workers.poll_idle_interval_s,
|
||||||
|
claim_timeout_s=ingester_cfg.workers.claim_timeout_s,
|
||||||
|
reaper_interval_s=ingester_cfg.workers.reaper_interval_s,
|
||||||
|
)
|
||||||
|
self._pollers = PollerManager(
|
||||||
|
configs=ingester_cfg.sources,
|
||||||
|
job_repo=self._jobs,
|
||||||
|
sync_repo=self._sync,
|
||||||
|
supported_extensions=supported_extensions,
|
||||||
|
default_max_attempts=ingester_cfg.workers.retry.max_attempts,
|
||||||
|
)
|
||||||
|
|
||||||
|
stop_event = asyncio.Event()
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
for sig in (signal.SIGINT, signal.SIGTERM):
|
||||||
|
try:
|
||||||
|
loop.add_signal_handler(sig, stop_event.set)
|
||||||
|
except NotImplementedError:
|
||||||
|
# Windows; signal handlers unavailable in asyncio.
|
||||||
|
pass
|
||||||
|
|
||||||
|
await self._pool.start()
|
||||||
|
await self._pollers.start()
|
||||||
|
logger.info(
|
||||||
|
"Ingester running: %d worker(s), %d source(s)",
|
||||||
|
ingester_cfg.workers.worker_count,
|
||||||
|
len(ingester_cfg.sources),
|
||||||
|
)
|
||||||
|
|
||||||
|
if api:
|
||||||
|
# HTTP control plane lands in a follow-up; for now this branch
|
||||||
|
# is a no-op so callers can still pass api=True without error.
|
||||||
|
logger.info(
|
||||||
|
"HTTP API not yet implemented; running pollers + workers only"
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
await stop_event.wait()
|
||||||
|
finally:
|
||||||
|
logger.info("Shutting down ingester")
|
||||||
|
await self._pollers.stop()
|
||||||
|
await self._pool.stop()
|
||||||
|
|
||||||
|
if self._queue_conn is not None:
|
||||||
|
await self._queue_conn.close()
|
||||||
|
self._queue_conn = None
|
||||||
|
|
@ -16,6 +16,7 @@ from haiku.rag.config import ( # noqa: E402
|
||||||
load_yaml_config,
|
load_yaml_config,
|
||||||
set_config,
|
set_config,
|
||||||
)
|
)
|
||||||
|
from haiku.rag.ingester.app import IngesterApp # noqa: E402
|
||||||
from haiku.rag.ingester.exceptions import PermanentError, TransientError # noqa: E402
|
from haiku.rag.ingester.exceptions import PermanentError, TransientError # noqa: E402
|
||||||
from haiku.rag.ingester.queue.migrations import open_queue # noqa: E402
|
from haiku.rag.ingester.queue.migrations import open_queue # noqa: E402
|
||||||
from haiku.rag.ingester.queue.models import Job, JobOp, JobStatus # noqa: E402
|
from haiku.rag.ingester.queue.models import Job, JobOp, JobStatus # noqa: E402
|
||||||
|
|
@ -37,7 +38,8 @@ cli.add_typer(queue_cli)
|
||||||
|
|
||||||
|
|
||||||
def _load_config_with_override(config_path: Path | None) -> AppConfig:
|
def _load_config_with_override(config_path: Path | None) -> AppConfig:
|
||||||
"""Mirror the haiku-rag CLI's config-loading pattern."""
|
"""Load AppConfig from `config_path`, the discovered project YAML, or the
|
||||||
|
process default — in that order."""
|
||||||
if config_path:
|
if config_path:
|
||||||
config = AppConfig.model_validate(load_yaml_config(config_path))
|
config = AppConfig.model_validate(load_yaml_config(config_path))
|
||||||
set_config(config)
|
set_config(config)
|
||||||
|
|
@ -96,6 +98,34 @@ def queue_migrate(
|
||||||
typer.echo(f"Queue at {path} is up to date")
|
typer.echo(f"Queue at {path} is up to date")
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_db_path(config: AppConfig, override: Path | None) -> Path:
|
||||||
|
return override or (config.storage.data_dir / "haiku.rag.lancedb")
|
||||||
|
|
||||||
|
|
||||||
|
@cli.command("serve")
|
||||||
|
def serve(
|
||||||
|
config: Path | None = typer.Option(
|
||||||
|
None, "--config", "-c", help="Path to haiku.rag.yaml."
|
||||||
|
),
|
||||||
|
db: Path | None = typer.Option(
|
||||||
|
None,
|
||||||
|
"--db",
|
||||||
|
help="LanceDB path (overrides config.storage.data_dir).",
|
||||||
|
),
|
||||||
|
no_api: bool = typer.Option(
|
||||||
|
False,
|
||||||
|
"--no-api",
|
||||||
|
help="Run pollers + workers without the HTTP control plane.",
|
||||||
|
),
|
||||||
|
) -> None:
|
||||||
|
"""Run the production ingester: pollers + workers (and the HTTP API
|
||||||
|
unless --no-api is set). Blocks until SIGINT/SIGTERM."""
|
||||||
|
app_config = _load_config_with_override(config)
|
||||||
|
db_path = _resolve_db_path(app_config, db)
|
||||||
|
app = IngesterApp(config=app_config, db_path=db_path)
|
||||||
|
asyncio.run(app.serve(api=not no_api))
|
||||||
|
|
||||||
|
|
||||||
@cli.command("run-once")
|
@cli.command("run-once")
|
||||||
def run_once(
|
def run_once(
|
||||||
uri: str = typer.Argument(..., help="URI to ingest (file://, http(s)://, s3://)."),
|
uri: str = typer.Argument(..., help="URI to ingest (file://, http(s)://, s3://)."),
|
||||||
|
|
@ -123,7 +153,7 @@ def run_once(
|
||||||
async def _run_once(
|
async def _run_once(
|
||||||
app_config: AppConfig, uri: str, db_path: Path | None, delete: bool
|
app_config: AppConfig, uri: str, db_path: Path | None, delete: bool
|
||||||
) -> None:
|
) -> None:
|
||||||
db = db_path or (app_config.storage.data_dir / "haiku.rag.lancedb")
|
db = _resolve_db_path(app_config, db_path)
|
||||||
now = datetime.now(UTC)
|
now = datetime.now(UTC)
|
||||||
job = Job(
|
job = Job(
|
||||||
id=f"adhoc-{uuid.uuid4()}",
|
id=f"adhoc-{uuid.uuid4()}",
|
||||||
|
|
|
||||||
13
haiku_rag_slim/haiku/rag/ingester/pollers/__init__.py
Normal file
13
haiku_rag_slim/haiku/rag/ingester/pollers/__init__.py
Normal file
|
|
@ -0,0 +1,13 @@
|
||||||
|
from haiku.rag.ingester.pollers.circuit_breaker import CircuitBreaker
|
||||||
|
from haiku.rag.ingester.pollers.factory import build_source
|
||||||
|
from haiku.rag.ingester.pollers.fs import FSPoller
|
||||||
|
from haiku.rag.ingester.pollers.manager import PollerManager
|
||||||
|
from haiku.rag.ingester.pollers.periodic import PeriodicPoller
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"CircuitBreaker",
|
||||||
|
"FSPoller",
|
||||||
|
"PeriodicPoller",
|
||||||
|
"PollerManager",
|
||||||
|
"build_source",
|
||||||
|
]
|
||||||
138
haiku_rag_slim/haiku/rag/ingester/pollers/base.py
Normal file
138
haiku_rag_slim/haiku/rag/ingester/pollers/base.py
Normal file
|
|
@ -0,0 +1,138 @@
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
|
from haiku.rag.config import SourceConfig
|
||||||
|
from haiku.rag.ingester.pollers.circuit_breaker import CircuitBreaker
|
||||||
|
from haiku.rag.ingester.queue.models import JobOp
|
||||||
|
from haiku.rag.ingester.queue.repository import JobRepo, SyncStateRepo
|
||||||
|
from haiku.rag.ingester.sources.base import (
|
||||||
|
Source,
|
||||||
|
SourceEvent,
|
||||||
|
SourceEventKind,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _enqueue_extra(cfg: SourceConfig) -> dict | None:
|
||||||
|
"""Per-source state worth carrying into the job (so the worker can rebuild
|
||||||
|
the same fetch context when it processes)."""
|
||||||
|
extra: dict = {}
|
||||||
|
storage_options = getattr(cfg, "storage_options", None)
|
||||||
|
if storage_options:
|
||||||
|
extra["storage_options"] = dict(storage_options)
|
||||||
|
headers = getattr(cfg, "headers", None)
|
||||||
|
if headers:
|
||||||
|
extra["headers"] = dict(headers)
|
||||||
|
return extra or None
|
||||||
|
|
||||||
|
|
||||||
|
def _max_attempts(cfg: SourceConfig, default: int) -> int:
|
||||||
|
return cfg.retry.max_attempts if cfg.retry is not None else default
|
||||||
|
|
||||||
|
|
||||||
|
class BasePoller:
|
||||||
|
"""Shared lifecycle: build a discover() coroutine + process its events
|
||||||
|
into queue jobs and sync_state updates. Subclasses provide the loop
|
||||||
|
(FS uses watchfiles + initial discover; periodic uses sleep+discover)."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
source: Source,
|
||||||
|
config: SourceConfig,
|
||||||
|
job_repo: JobRepo,
|
||||||
|
sync_repo: SyncStateRepo,
|
||||||
|
breaker: CircuitBreaker | None = None,
|
||||||
|
default_max_attempts: int = 5,
|
||||||
|
):
|
||||||
|
self.source = source
|
||||||
|
self.config = config
|
||||||
|
self._jobs = job_repo
|
||||||
|
self._sync = sync_repo
|
||||||
|
self._breaker = breaker or CircuitBreaker(config.circuit_breaker)
|
||||||
|
self._stop = asyncio.Event()
|
||||||
|
self._task: asyncio.Task | None = None
|
||||||
|
self._last_polled_at: datetime | None = None
|
||||||
|
self._default_max_attempts = default_max_attempts
|
||||||
|
|
||||||
|
@property
|
||||||
|
def source_id(self) -> str:
|
||||||
|
return self.source.source_id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def last_polled_at(self) -> datetime | None:
|
||||||
|
return self._last_polled_at
|
||||||
|
|
||||||
|
async def run(self) -> None: # pragma: no cover - subclasses override
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
self._stop.set()
|
||||||
|
if self._task is not None:
|
||||||
|
await asyncio.gather(self._task, return_exceptions=True)
|
||||||
|
self._task = None
|
||||||
|
|
||||||
|
async def _sweep_once(self) -> bool:
|
||||||
|
"""One discover() sweep. Returns True on success, False if the
|
||||||
|
breaker is open or the sweep failed (and was recorded)."""
|
||||||
|
if self._breaker.is_open:
|
||||||
|
logger.debug(
|
||||||
|
"Skipping discover() — circuit breaker open for %s", self.source_id
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
snapshot = await self._sync.get_snapshot(self.source_id)
|
||||||
|
async for event in self.source.discover(since=snapshot):
|
||||||
|
await self._handle_event(event)
|
||||||
|
self._breaker.record_success()
|
||||||
|
self._last_polled_at = datetime.now(UTC)
|
||||||
|
return True
|
||||||
|
except Exception as exc:
|
||||||
|
self._breaker.record_failure()
|
||||||
|
logger.exception(
|
||||||
|
"discover() failed for %s (consecutive=%d): %s",
|
||||||
|
self.source_id,
|
||||||
|
self._breaker.consecutive_failures,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _handle_event(self, event: SourceEvent) -> None:
|
||||||
|
if event.kind is SourceEventKind.UPSERT:
|
||||||
|
await self._jobs.enqueue(
|
||||||
|
event.source_id,
|
||||||
|
event.uri,
|
||||||
|
op=JobOp.UPSERT,
|
||||||
|
revision=event.revision,
|
||||||
|
max_attempts=_max_attempts(self.config, self._default_max_attempts),
|
||||||
|
extra=_enqueue_extra(self.config),
|
||||||
|
)
|
||||||
|
# Don't write revision to sync_state here — the worker writes it
|
||||||
|
# after a successful ingestion. last_seen_at gets bumped to keep
|
||||||
|
# orphan detection accurate.
|
||||||
|
await self._sync.upsert(
|
||||||
|
event.source_id,
|
||||||
|
event.uri,
|
||||||
|
revision=None,
|
||||||
|
content_hash=None,
|
||||||
|
)
|
||||||
|
elif event.kind is SourceEventKind.UNCHANGED:
|
||||||
|
# Touch last_seen_at without changing the stored revision.
|
||||||
|
await self._sync.upsert(
|
||||||
|
event.source_id,
|
||||||
|
event.uri,
|
||||||
|
revision=event.revision,
|
||||||
|
content_hash=None,
|
||||||
|
)
|
||||||
|
elif event.kind is SourceEventKind.DELETE:
|
||||||
|
if not self.config.delete_orphans:
|
||||||
|
return
|
||||||
|
await self._jobs.enqueue(
|
||||||
|
event.source_id,
|
||||||
|
event.uri,
|
||||||
|
op=JobOp.DELETE,
|
||||||
|
max_attempts=_max_attempts(self.config, self._default_max_attempts),
|
||||||
|
extra=_enqueue_extra(self.config),
|
||||||
|
)
|
||||||
48
haiku_rag_slim/haiku/rag/ingester/pollers/circuit_breaker.py
Normal file
48
haiku_rag_slim/haiku/rag/ingester/pollers/circuit_breaker.py
Normal file
|
|
@ -0,0 +1,48 @@
|
||||||
|
import time
|
||||||
|
from collections.abc import Callable
|
||||||
|
|
||||||
|
from haiku.rag.config import CircuitBreakerConfig
|
||||||
|
|
||||||
|
|
||||||
|
class CircuitBreaker:
|
||||||
|
"""Three-state breaker over discover() failures.
|
||||||
|
|
||||||
|
- closed: failures are counted; threshold flips to open.
|
||||||
|
- open: probes are blocked until cooldown elapses, then a single probe
|
||||||
|
is allowed; success closes the breaker, another failure re-opens it.
|
||||||
|
|
||||||
|
`now_fn` is injectable so tests don't need monkeypatching of time.time.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: CircuitBreakerConfig | None = None,
|
||||||
|
*,
|
||||||
|
now_fn: Callable[[], float] = time.monotonic,
|
||||||
|
):
|
||||||
|
self._config = config or CircuitBreakerConfig()
|
||||||
|
self._now = now_fn
|
||||||
|
self._consecutive_failures = 0
|
||||||
|
self._opened_at: float | None = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_open(self) -> bool:
|
||||||
|
if self._opened_at is None:
|
||||||
|
return False
|
||||||
|
if self._now() - self._opened_at >= self._config.cooldown_s:
|
||||||
|
# cooldown elapsed; let the next call probe
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
def record_success(self) -> None:
|
||||||
|
self._consecutive_failures = 0
|
||||||
|
self._opened_at = None
|
||||||
|
|
||||||
|
def record_failure(self) -> None:
|
||||||
|
self._consecutive_failures += 1
|
||||||
|
if self._consecutive_failures >= self._config.failure_threshold:
|
||||||
|
self._opened_at = self._now()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def consecutive_failures(self) -> int:
|
||||||
|
return self._consecutive_failures
|
||||||
41
haiku_rag_slim/haiku/rag/ingester/pollers/factory.py
Normal file
41
haiku_rag_slim/haiku/rag/ingester/pollers/factory.py
Normal file
|
|
@ -0,0 +1,41 @@
|
||||||
|
from haiku.rag.config import (
|
||||||
|
FSSourceConfig,
|
||||||
|
HTTPSourceConfig,
|
||||||
|
S3SourceConfig,
|
||||||
|
SourceConfig,
|
||||||
|
)
|
||||||
|
from haiku.rag.ingester.sources import FSSource, HTTPSource, S3Source, Source
|
||||||
|
|
||||||
|
|
||||||
|
def build_source(
|
||||||
|
cfg: SourceConfig,
|
||||||
|
*,
|
||||||
|
supported_extensions: list[str] | None = None,
|
||||||
|
) -> Source:
|
||||||
|
"""Instantiate the right adapter for a SourceConfig.
|
||||||
|
|
||||||
|
Source IDs auto-derive from the target when the config didn't supply one,
|
||||||
|
matching the conventions in the adapters themselves (fs:<root>,
|
||||||
|
s3:<bucket>/<prefix>, http:<id>).
|
||||||
|
"""
|
||||||
|
if isinstance(cfg, FSSourceConfig):
|
||||||
|
return FSSource(
|
||||||
|
root=cfg.root,
|
||||||
|
ignore_patterns=cfg.ignore_patterns or None,
|
||||||
|
include_patterns=cfg.include_patterns or None,
|
||||||
|
supported_extensions=supported_extensions,
|
||||||
|
)
|
||||||
|
if isinstance(cfg, HTTPSourceConfig):
|
||||||
|
if cfg.id is None:
|
||||||
|
raise ValueError("HTTPSourceConfig.id is required")
|
||||||
|
return HTTPSource(source_id=cfg.id, urls=cfg.urls, headers=cfg.headers)
|
||||||
|
if isinstance(cfg, S3SourceConfig):
|
||||||
|
return S3Source(
|
||||||
|
uri=cfg.uri,
|
||||||
|
storage_options=cfg.storage_options,
|
||||||
|
ignore_patterns=cfg.ignore_patterns or None,
|
||||||
|
include_patterns=cfg.include_patterns or None,
|
||||||
|
supported_extensions=supported_extensions,
|
||||||
|
source_id=cfg.id,
|
||||||
|
)
|
||||||
|
raise TypeError(f"Unsupported source config: {type(cfg).__name__}")
|
||||||
124
haiku_rag_slim/haiku/rag/ingester/pollers/fs.py
Normal file
124
haiku_rag_slim/haiku/rag/ingester/pollers/fs.py
Normal file
|
|
@ -0,0 +1,124 @@
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from watchfiles import Change, awatch
|
||||||
|
|
||||||
|
from haiku.rag.ingester.pollers.base import BasePoller
|
||||||
|
from haiku.rag.ingester.queue.models import JobOp
|
||||||
|
from haiku.rag.ingester.sources.filter import FileFilter
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from haiku.rag.config import FSSourceConfig
|
||||||
|
from haiku.rag.ingester.pollers.circuit_breaker import CircuitBreaker
|
||||||
|
from haiku.rag.ingester.sources.fs import FSSource
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class FSPoller(BasePoller):
|
||||||
|
"""Filesystem poller: initial discover() sweep plus a watchfiles-driven
|
||||||
|
push loop. Periodic sweeps still run so files modified while the watcher
|
||||||
|
was offline get picked up too."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
source: "FSSource",
|
||||||
|
config: "FSSourceConfig",
|
||||||
|
job_repo,
|
||||||
|
sync_repo,
|
||||||
|
breaker: "CircuitBreaker | None" = None,
|
||||||
|
default_max_attempts: int = 5,
|
||||||
|
):
|
||||||
|
super().__init__(
|
||||||
|
source=source,
|
||||||
|
config=config,
|
||||||
|
job_repo=job_repo,
|
||||||
|
sync_repo=sync_repo,
|
||||||
|
breaker=breaker,
|
||||||
|
default_max_attempts=default_max_attempts,
|
||||||
|
)
|
||||||
|
self._fs_source: FSSource = source
|
||||||
|
self._fs_config: FSSourceConfig = config
|
||||||
|
self._filter = FileFilter(
|
||||||
|
ignore_patterns=config.ignore_patterns or None,
|
||||||
|
include_patterns=config.include_patterns or None,
|
||||||
|
supported_extensions=source.supported_extensions,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def run(self) -> None:
|
||||||
|
await self._sweep_once()
|
||||||
|
watch_task = asyncio.create_task(self._watch_loop())
|
||||||
|
sweep_task = asyncio.create_task(self._sweep_loop())
|
||||||
|
try:
|
||||||
|
await self._stop.wait()
|
||||||
|
finally:
|
||||||
|
watch_task.cancel()
|
||||||
|
sweep_task.cancel()
|
||||||
|
await asyncio.gather(watch_task, sweep_task, return_exceptions=True)
|
||||||
|
|
||||||
|
async def _sweep_loop(self) -> None:
|
||||||
|
"""Periodic full sweep. Catches files modified while the watcher
|
||||||
|
wasn't running (gaps between starts, races, FS events the OS dropped)."""
|
||||||
|
while not self._stop.is_set():
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(
|
||||||
|
self._stop.wait(), timeout=self.config.poll_interval_s
|
||||||
|
)
|
||||||
|
return
|
||||||
|
except TimeoutError:
|
||||||
|
pass
|
||||||
|
await self._sweep_once()
|
||||||
|
|
||||||
|
async def _watch_loop(self) -> None:
|
||||||
|
"""Push-event loop on top of watchfiles. Each change is translated
|
||||||
|
into one queue job — no need to re-stat or re-snapshot."""
|
||||||
|
try:
|
||||||
|
async for changes in awatch(
|
||||||
|
self._fs_source.root,
|
||||||
|
watch_filter=self._filter,
|
||||||
|
stop_event=self._stop,
|
||||||
|
):
|
||||||
|
for change, path in changes:
|
||||||
|
await self._handle_watch_change(change, Path(path))
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
self._breaker.record_failure()
|
||||||
|
logger.exception("watchfiles loop failed for %s: %s", self.source_id, exc)
|
||||||
|
|
||||||
|
async def _handle_watch_change(self, change: Change, path: Path) -> None:
|
||||||
|
uri = path.as_uri()
|
||||||
|
if change is Change.deleted:
|
||||||
|
if not self._fs_config.delete_orphans:
|
||||||
|
return
|
||||||
|
await self._jobs.enqueue(
|
||||||
|
self.source_id,
|
||||||
|
uri,
|
||||||
|
op=JobOp.DELETE,
|
||||||
|
max_attempts=self._max_attempts(),
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if change in (Change.added, Change.modified):
|
||||||
|
revision = str(path.stat().st_mtime_ns) if path.exists() else None
|
||||||
|
await self._jobs.enqueue(
|
||||||
|
self.source_id,
|
||||||
|
uri,
|
||||||
|
op=JobOp.UPSERT,
|
||||||
|
revision=revision,
|
||||||
|
max_attempts=self._max_attempts(),
|
||||||
|
)
|
||||||
|
await self._sync.upsert(
|
||||||
|
self.source_id, uri, revision=None, content_hash=None
|
||||||
|
)
|
||||||
|
|
||||||
|
def _max_attempts(self) -> int:
|
||||||
|
cfg = self._fs_config
|
||||||
|
return (
|
||||||
|
cfg.retry.max_attempts
|
||||||
|
if cfg.retry is not None
|
||||||
|
else self._default_max_attempts
|
||||||
|
)
|
||||||
89
haiku_rag_slim/haiku/rag/ingester/pollers/manager.py
Normal file
89
haiku_rag_slim/haiku/rag/ingester/pollers/manager.py
Normal file
|
|
@ -0,0 +1,89 @@
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from haiku.rag.config import FSSourceConfig, SourceConfig
|
||||||
|
from haiku.rag.ingester.pollers.base import BasePoller
|
||||||
|
from haiku.rag.ingester.pollers.circuit_breaker import CircuitBreaker
|
||||||
|
from haiku.rag.ingester.pollers.factory import build_source
|
||||||
|
from haiku.rag.ingester.pollers.fs import FSPoller
|
||||||
|
from haiku.rag.ingester.pollers.periodic import PeriodicPoller
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from haiku.rag.ingester.queue.repository import JobRepo, SyncStateRepo
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class PollerManager:
|
||||||
|
"""Owns one poller per configured source. Lifecycle: build → start →
|
||||||
|
stop. Each poller runs as an independent asyncio task; failures in one
|
||||||
|
don't affect the others."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
configs: list[SourceConfig],
|
||||||
|
job_repo: "JobRepo",
|
||||||
|
sync_repo: "SyncStateRepo",
|
||||||
|
supported_extensions: list[str] | None = None,
|
||||||
|
default_max_attempts: int = 5,
|
||||||
|
):
|
||||||
|
self._configs = configs
|
||||||
|
self._jobs = job_repo
|
||||||
|
self._sync = sync_repo
|
||||||
|
self._supported_extensions = supported_extensions
|
||||||
|
self._default_max_attempts = default_max_attempts
|
||||||
|
self._pollers: list[BasePoller] = []
|
||||||
|
self._tasks: list[asyncio.Task] = []
|
||||||
|
|
||||||
|
def build_pollers(self) -> list[BasePoller]:
|
||||||
|
pollers: list[BasePoller] = []
|
||||||
|
for cfg in self._configs:
|
||||||
|
source = build_source(cfg, supported_extensions=self._supported_extensions)
|
||||||
|
breaker = CircuitBreaker(cfg.circuit_breaker)
|
||||||
|
if isinstance(cfg, FSSourceConfig):
|
||||||
|
from haiku.rag.ingester.sources.fs import FSSource
|
||||||
|
|
||||||
|
assert isinstance(source, FSSource)
|
||||||
|
pollers.append(
|
||||||
|
FSPoller(
|
||||||
|
source=source,
|
||||||
|
config=cfg,
|
||||||
|
job_repo=self._jobs,
|
||||||
|
sync_repo=self._sync,
|
||||||
|
breaker=breaker,
|
||||||
|
default_max_attempts=self._default_max_attempts,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
pollers.append(
|
||||||
|
PeriodicPoller(
|
||||||
|
source=source,
|
||||||
|
config=cfg,
|
||||||
|
job_repo=self._jobs,
|
||||||
|
sync_repo=self._sync,
|
||||||
|
breaker=breaker,
|
||||||
|
default_max_attempts=self._default_max_attempts,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return pollers
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
if self._pollers:
|
||||||
|
raise RuntimeError("PollerManager already started")
|
||||||
|
self._pollers = self.build_pollers()
|
||||||
|
for poller in self._pollers:
|
||||||
|
self._tasks.append(asyncio.create_task(poller.run()))
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
for poller in self._pollers:
|
||||||
|
await poller.stop()
|
||||||
|
if self._tasks:
|
||||||
|
await asyncio.gather(*self._tasks, return_exceptions=True)
|
||||||
|
self._tasks.clear()
|
||||||
|
self._pollers.clear()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def pollers(self) -> list[BasePoller]:
|
||||||
|
return list(self._pollers)
|
||||||
48
haiku_rag_slim/haiku/rag/ingester/pollers/periodic.py
Normal file
48
haiku_rag_slim/haiku/rag/ingester/pollers/periodic.py
Normal file
|
|
@ -0,0 +1,48 @@
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from haiku.rag.ingester.pollers.base import BasePoller
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from haiku.rag.ingester.pollers.circuit_breaker import CircuitBreaker
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class PeriodicPoller(BasePoller):
|
||||||
|
"""Runs `source.discover()` on a fixed interval. Used for HTTP, S3, WebDAV
|
||||||
|
— sources that only know about changes when we ask them."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
source,
|
||||||
|
config,
|
||||||
|
job_repo,
|
||||||
|
sync_repo,
|
||||||
|
breaker: "CircuitBreaker | None" = None,
|
||||||
|
default_max_attempts: int = 5,
|
||||||
|
):
|
||||||
|
super().__init__(
|
||||||
|
source=source,
|
||||||
|
config=config,
|
||||||
|
job_repo=job_repo,
|
||||||
|
sync_repo=sync_repo,
|
||||||
|
breaker=breaker,
|
||||||
|
default_max_attempts=default_max_attempts,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def run(self) -> None:
|
||||||
|
# Initial sweep on startup so newly-configured sources are scanned
|
||||||
|
# immediately instead of waiting one full interval.
|
||||||
|
await self._sweep_once()
|
||||||
|
while not self._stop.is_set():
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(
|
||||||
|
self._stop.wait(), timeout=self.config.poll_interval_s
|
||||||
|
)
|
||||||
|
return
|
||||||
|
except TimeoutError:
|
||||||
|
pass
|
||||||
|
await self._sweep_once()
|
||||||
|
|
@ -140,8 +140,8 @@ class S3Source:
|
||||||
discovered_at=now,
|
discovered_at=now,
|
||||||
)
|
)
|
||||||
|
|
||||||
# URIs we previously synced but that no longer appear under the
|
# URIs in the snapshot that no longer appear under the prefix have
|
||||||
# prefix have been deleted upstream.
|
# been deleted upstream — emit DELETE so the poller cleans up.
|
||||||
for uri in snapshot:
|
for uri in snapshot:
|
||||||
if uri in seen:
|
if uri in seen:
|
||||||
continue
|
continue
|
||||||
|
|
|
||||||
82
tests/ingester/test_circuit_breaker.py
Normal file
82
tests/ingester/test_circuit_breaker.py
Normal file
|
|
@ -0,0 +1,82 @@
|
||||||
|
from haiku.rag.config import CircuitBreakerConfig
|
||||||
|
from haiku.rag.ingester.pollers.circuit_breaker import CircuitBreaker
|
||||||
|
|
||||||
|
|
||||||
|
class _Clock:
|
||||||
|
def __init__(self, start: float = 0.0):
|
||||||
|
self.now = start
|
||||||
|
|
||||||
|
def __call__(self) -> float:
|
||||||
|
return self.now
|
||||||
|
|
||||||
|
def advance(self, seconds: float) -> None:
|
||||||
|
self.now += seconds
|
||||||
|
|
||||||
|
|
||||||
|
def test_starts_closed():
|
||||||
|
breaker = CircuitBreaker(CircuitBreakerConfig(failure_threshold=3))
|
||||||
|
assert breaker.is_open is False
|
||||||
|
assert breaker.consecutive_failures == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_opens_after_threshold():
|
||||||
|
clock = _Clock()
|
||||||
|
breaker = CircuitBreaker(
|
||||||
|
CircuitBreakerConfig(failure_threshold=3, cooldown_s=60.0), now_fn=clock
|
||||||
|
)
|
||||||
|
breaker.record_failure()
|
||||||
|
breaker.record_failure()
|
||||||
|
assert breaker.is_open is False # threshold not reached
|
||||||
|
breaker.record_failure()
|
||||||
|
assert breaker.is_open is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_success_resets_failure_count():
|
||||||
|
breaker = CircuitBreaker(CircuitBreakerConfig(failure_threshold=3, cooldown_s=60.0))
|
||||||
|
breaker.record_failure()
|
||||||
|
breaker.record_failure()
|
||||||
|
breaker.record_success()
|
||||||
|
assert breaker.consecutive_failures == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_cooldown_allows_probe():
|
||||||
|
clock = _Clock()
|
||||||
|
breaker = CircuitBreaker(
|
||||||
|
CircuitBreakerConfig(failure_threshold=2, cooldown_s=10.0), now_fn=clock
|
||||||
|
)
|
||||||
|
breaker.record_failure()
|
||||||
|
breaker.record_failure()
|
||||||
|
assert breaker.is_open is True
|
||||||
|
|
||||||
|
clock.advance(5.0)
|
||||||
|
assert breaker.is_open is True # still cooling
|
||||||
|
clock.advance(5.5)
|
||||||
|
assert breaker.is_open is False # cooldown elapsed → probe allowed
|
||||||
|
|
||||||
|
|
||||||
|
def test_probe_failure_reopens():
|
||||||
|
clock = _Clock()
|
||||||
|
breaker = CircuitBreaker(
|
||||||
|
CircuitBreakerConfig(failure_threshold=2, cooldown_s=10.0), now_fn=clock
|
||||||
|
)
|
||||||
|
breaker.record_failure()
|
||||||
|
breaker.record_failure()
|
||||||
|
clock.advance(15.0)
|
||||||
|
assert breaker.is_open is False
|
||||||
|
|
||||||
|
breaker.record_failure()
|
||||||
|
# failure_threshold=2 already exceeded by accumulating — breaker re-opens
|
||||||
|
assert breaker.is_open is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_probe_success_closes():
|
||||||
|
clock = _Clock()
|
||||||
|
breaker = CircuitBreaker(
|
||||||
|
CircuitBreakerConfig(failure_threshold=2, cooldown_s=10.0), now_fn=clock
|
||||||
|
)
|
||||||
|
breaker.record_failure()
|
||||||
|
breaker.record_failure()
|
||||||
|
clock.advance(15.0)
|
||||||
|
breaker.record_success()
|
||||||
|
assert breaker.is_open is False
|
||||||
|
assert breaker.consecutive_failures == 0
|
||||||
111
tests/ingester/test_config.py
Normal file
111
tests/ingester/test_config.py
Normal file
|
|
@ -0,0 +1,111 @@
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import yaml
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
from haiku.rag.config import (
|
||||||
|
AppConfig,
|
||||||
|
FSSourceConfig,
|
||||||
|
HTTPSourceConfig,
|
||||||
|
IngesterConfig,
|
||||||
|
RetryPolicyConfig,
|
||||||
|
S3SourceConfig,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_default_ingester_config_has_sane_values():
|
||||||
|
cfg = IngesterConfig()
|
||||||
|
assert cfg.sources == []
|
||||||
|
assert cfg.workers.worker_count == 4
|
||||||
|
assert cfg.workers.retry.max_attempts == 5
|
||||||
|
assert cfg.api.enabled is True
|
||||||
|
assert cfg.api.port == 8765
|
||||||
|
|
||||||
|
|
||||||
|
def test_discriminator_picks_fs_source():
|
||||||
|
cfg = IngesterConfig.model_validate(
|
||||||
|
{"sources": [{"type": "fs", "root": "/tmp/docs"}]}
|
||||||
|
)
|
||||||
|
assert isinstance(cfg.sources[0], FSSourceConfig)
|
||||||
|
assert cfg.sources[0].root == Path("/tmp/docs")
|
||||||
|
assert cfg.sources[0].delete_orphans is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_discriminator_picks_http_source():
|
||||||
|
cfg = IngesterConfig.model_validate(
|
||||||
|
{"sources": [{"type": "http", "id": "arxiv", "urls": ["https://x"]}]}
|
||||||
|
)
|
||||||
|
assert isinstance(cfg.sources[0], HTTPSourceConfig)
|
||||||
|
assert cfg.sources[0].id == "arxiv"
|
||||||
|
|
||||||
|
|
||||||
|
def test_discriminator_picks_s3_source():
|
||||||
|
cfg = IngesterConfig.model_validate(
|
||||||
|
{"sources": [{"type": "s3", "uri": "s3://bucket/prefix/"}]}
|
||||||
|
)
|
||||||
|
assert isinstance(cfg.sources[0], S3SourceConfig)
|
||||||
|
assert cfg.sources[0].uri == "s3://bucket/prefix/"
|
||||||
|
|
||||||
|
|
||||||
|
def test_discriminator_rejects_unknown_type():
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
IngesterConfig.model_validate({"sources": [{"type": "ftp", "uri": "x"}]})
|
||||||
|
|
||||||
|
|
||||||
|
def test_per_source_retry_overrides_default():
|
||||||
|
cfg = IngesterConfig.model_validate(
|
||||||
|
{
|
||||||
|
"sources": [
|
||||||
|
{
|
||||||
|
"type": "fs",
|
||||||
|
"root": "/tmp",
|
||||||
|
"retry": {"max_attempts": 10},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
)
|
||||||
|
src = cfg.sources[0]
|
||||||
|
assert src.retry is not None
|
||||||
|
assert src.retry.max_attempts == 10
|
||||||
|
# other retry fields fall back to the RetryPolicyConfig defaults
|
||||||
|
assert src.retry.base_delay_s == RetryPolicyConfig().base_delay_s
|
||||||
|
|
||||||
|
|
||||||
|
def test_yaml_round_trip():
|
||||||
|
yaml_text = """
|
||||||
|
ingester:
|
||||||
|
sources:
|
||||||
|
- type: fs
|
||||||
|
root: /data/docs
|
||||||
|
ignore_patterns: ["**/.git/**"]
|
||||||
|
delete_orphans: true
|
||||||
|
- type: s3
|
||||||
|
uri: s3://my-bucket/incoming/
|
||||||
|
poll_interval_s: 300
|
||||||
|
storage_options:
|
||||||
|
endpoint: http://seaweed:8333
|
||||||
|
- type: http
|
||||||
|
id: arxiv
|
||||||
|
urls: [https://arxiv.org/pdf/2301.12345.pdf]
|
||||||
|
headers:
|
||||||
|
Authorization: Bearer abc
|
||||||
|
poll_interval_s: 86400
|
||||||
|
workers:
|
||||||
|
worker_count: 8
|
||||||
|
max_concurrent: 4
|
||||||
|
api:
|
||||||
|
enabled: false
|
||||||
|
"""
|
||||||
|
data = yaml.safe_load(yaml_text)
|
||||||
|
app = AppConfig.model_validate(data)
|
||||||
|
assert len(app.ingester.sources) == 3
|
||||||
|
fs, s3, http = app.ingester.sources
|
||||||
|
assert isinstance(fs, FSSourceConfig)
|
||||||
|
assert isinstance(s3, S3SourceConfig)
|
||||||
|
assert isinstance(http, HTTPSourceConfig)
|
||||||
|
assert fs.ignore_patterns == ["**/.git/**"]
|
||||||
|
assert s3.storage_options["endpoint"] == "http://seaweed:8333"
|
||||||
|
assert http.headers["Authorization"] == "Bearer abc"
|
||||||
|
assert app.ingester.workers.worker_count == 8
|
||||||
|
assert app.ingester.api.enabled is False
|
||||||
338
tests/ingester/test_pollers.py
Normal file
338
tests/ingester/test_pollers.py
Normal file
|
|
@ -0,0 +1,338 @@
|
||||||
|
import asyncio
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import aiosqlite
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from haiku.rag.config import (
|
||||||
|
CircuitBreakerConfig,
|
||||||
|
FSSourceConfig,
|
||||||
|
HTTPSourceConfig,
|
||||||
|
S3SourceConfig,
|
||||||
|
)
|
||||||
|
from haiku.rag.ingester.pollers.circuit_breaker import CircuitBreaker
|
||||||
|
from haiku.rag.ingester.pollers.manager import PollerManager
|
||||||
|
from haiku.rag.ingester.pollers.periodic import PeriodicPoller
|
||||||
|
from haiku.rag.ingester.queue.migrations import apply_migrations
|
||||||
|
from haiku.rag.ingester.queue.models import JobOp, JobStatus
|
||||||
|
from haiku.rag.ingester.queue.repository import JobRepo, SyncStateRepo
|
||||||
|
from haiku.rag.ingester.sources.base import (
|
||||||
|
FetchResult,
|
||||||
|
SourceEvent,
|
||||||
|
SourceEventKind,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
async def conn(tmp_path):
|
||||||
|
path = tmp_path / "queue.db"
|
||||||
|
connection = await aiosqlite.connect(str(path))
|
||||||
|
connection.row_factory = aiosqlite.Row
|
||||||
|
await apply_migrations(connection)
|
||||||
|
yield connection
|
||||||
|
await connection.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def jobs(conn):
|
||||||
|
return JobRepo(conn)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def sync(conn):
|
||||||
|
return SyncStateRepo(conn)
|
||||||
|
|
||||||
|
|
||||||
|
class _StubSource:
|
||||||
|
"""Test double that yields a scripted sequence of events on each
|
||||||
|
discover() call. `fetch` and `supports` aren't exercised by pollers."""
|
||||||
|
|
||||||
|
def __init__(self, source_id: str, sweeps: list[list[SourceEvent]]):
|
||||||
|
self.source_id = source_id
|
||||||
|
self._sweeps = list(sweeps)
|
||||||
|
self.discover_calls = 0
|
||||||
|
self.fail_with: Exception | None = None
|
||||||
|
|
||||||
|
def supports(self, uri: str) -> bool: # pragma: no cover - unused here
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def head(self, uri: str) -> str | None: # pragma: no cover
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def fetch(self, uri: str) -> FetchResult: # pragma: no cover
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
async def discover(self, since=None):
|
||||||
|
self.discover_calls += 1
|
||||||
|
if self.fail_with is not None:
|
||||||
|
raise self.fail_with
|
||||||
|
events = self._sweeps.pop(0) if self._sweeps else []
|
||||||
|
for event in events:
|
||||||
|
yield event
|
||||||
|
|
||||||
|
|
||||||
|
def _event(
|
||||||
|
uri: str,
|
||||||
|
kind=SourceEventKind.UPSERT,
|
||||||
|
revision: str | None = "v1",
|
||||||
|
source_id: str = "src",
|
||||||
|
):
|
||||||
|
return SourceEvent(
|
||||||
|
source_id=source_id,
|
||||||
|
uri=uri,
|
||||||
|
kind=kind,
|
||||||
|
revision=None if kind is SourceEventKind.DELETE else revision,
|
||||||
|
discovered_at=datetime.now(UTC),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def fs_config(tmp_path):
|
||||||
|
return FSSourceConfig(
|
||||||
|
type="fs",
|
||||||
|
id="src",
|
||||||
|
root=tmp_path,
|
||||||
|
delete_orphans=True,
|
||||||
|
poll_interval_s=0.05,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _periodic(source, config, jobs, sync, **kwargs):
|
||||||
|
return PeriodicPoller(
|
||||||
|
source=source,
|
||||||
|
config=config,
|
||||||
|
job_repo=jobs,
|
||||||
|
sync_repo=sync,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# --- _sweep_once / event handling on the base class via PeriodicPoller ---
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_upsert_event_enqueues_job_and_touches_sync_state(fs_config, jobs, sync):
|
||||||
|
source = _StubSource("src", [[_event("file:///a.md", revision="r1")]])
|
||||||
|
poller = _periodic(source, fs_config, jobs, sync)
|
||||||
|
ok = await poller._sweep_once()
|
||||||
|
assert ok is True
|
||||||
|
|
||||||
|
queued = await jobs.list_jobs(source_id="src")
|
||||||
|
assert len(queued) == 1
|
||||||
|
assert queued[0].op is JobOp.UPSERT
|
||||||
|
assert queued[0].revision == "r1"
|
||||||
|
|
||||||
|
# Pollers DO NOT write revision to sync_state — the worker does that
|
||||||
|
# after a successful ingest. But last_seen_at is bumped.
|
||||||
|
snapshot = await sync.get_snapshot("src")
|
||||||
|
assert snapshot == {} # revision left empty by the poller
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_unchanged_event_touches_sync_state_no_job(fs_config, jobs, sync):
|
||||||
|
source = _StubSource(
|
||||||
|
"src", [[_event("file:///a.md", kind=SourceEventKind.UNCHANGED, revision="r1")]]
|
||||||
|
)
|
||||||
|
poller = _periodic(source, fs_config, jobs, sync)
|
||||||
|
await poller._sweep_once()
|
||||||
|
|
||||||
|
assert await jobs.list_jobs(source_id="src") == []
|
||||||
|
assert await sync.get_snapshot("src") == {"file:///a.md": "r1"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_delete_event_enqueues_delete_job(fs_config, jobs, sync):
|
||||||
|
source = _StubSource(
|
||||||
|
"src", [[_event("file:///gone.md", kind=SourceEventKind.DELETE)]]
|
||||||
|
)
|
||||||
|
poller = _periodic(source, fs_config, jobs, sync)
|
||||||
|
await poller._sweep_once()
|
||||||
|
|
||||||
|
queued = await jobs.list_jobs(source_id="src")
|
||||||
|
assert len(queued) == 1
|
||||||
|
assert queued[0].op is JobOp.DELETE
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_delete_event_skipped_when_delete_orphans_false(fs_config, jobs, sync):
|
||||||
|
fs_config = fs_config.model_copy(update={"delete_orphans": False})
|
||||||
|
source = _StubSource(
|
||||||
|
"src", [[_event("file:///gone.md", kind=SourceEventKind.DELETE)]]
|
||||||
|
)
|
||||||
|
poller = _periodic(source, fs_config, jobs, sync)
|
||||||
|
await poller._sweep_once()
|
||||||
|
assert await jobs.list_jobs(source_id="src") == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_repeated_sweep_does_not_duplicate_jobs(fs_config, jobs, sync):
|
||||||
|
"""Live-uniqueness: enqueueing the same (source_id, uri, op) when a job
|
||||||
|
is already queued/claimed is a no-op."""
|
||||||
|
event = _event("file:///a.md", revision="r1")
|
||||||
|
source = _StubSource("src", [[event], [event]])
|
||||||
|
poller = _periodic(source, fs_config, jobs, sync)
|
||||||
|
await poller._sweep_once()
|
||||||
|
await poller._sweep_once()
|
||||||
|
queued = await jobs.list_jobs(source_id="src")
|
||||||
|
assert len(queued) == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_circuit_breaker_pauses_sweeps_after_failures(fs_config, jobs, sync):
|
||||||
|
class _Clock:
|
||||||
|
now = 0.0
|
||||||
|
|
||||||
|
def __call__(self):
|
||||||
|
return self.now
|
||||||
|
|
||||||
|
clock = _Clock()
|
||||||
|
breaker = CircuitBreaker(
|
||||||
|
CircuitBreakerConfig(failure_threshold=2, cooldown_s=30.0),
|
||||||
|
now_fn=clock,
|
||||||
|
)
|
||||||
|
source = _StubSource("src", [])
|
||||||
|
source.fail_with = RuntimeError("upstream down")
|
||||||
|
|
||||||
|
poller = _periodic(source, fs_config, jobs, sync, breaker=breaker)
|
||||||
|
# Two failures open the breaker.
|
||||||
|
assert await poller._sweep_once() is False
|
||||||
|
assert await poller._sweep_once() is False
|
||||||
|
assert breaker.is_open is True
|
||||||
|
|
||||||
|
# Third call should be skipped — discover() is not invoked.
|
||||||
|
before = source.discover_calls
|
||||||
|
assert await poller._sweep_once() is False
|
||||||
|
assert source.discover_calls == before
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_sweep_records_last_polled_at_on_success(fs_config, jobs, sync):
|
||||||
|
source = _StubSource("src", [[]])
|
||||||
|
poller = _periodic(source, fs_config, jobs, sync)
|
||||||
|
assert poller.last_polled_at is None
|
||||||
|
await poller._sweep_once()
|
||||||
|
assert poller.last_polled_at is not None
|
||||||
|
assert poller.last_polled_at.tzinfo is not None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_per_source_retry_policy_overrides_default(jobs, sync, tmp_path):
|
||||||
|
from haiku.rag.config import RetryPolicyConfig
|
||||||
|
|
||||||
|
cfg = FSSourceConfig(
|
||||||
|
type="fs",
|
||||||
|
id="src",
|
||||||
|
root=tmp_path,
|
||||||
|
retry=RetryPolicyConfig(max_attempts=9),
|
||||||
|
)
|
||||||
|
source = _StubSource("src", [[_event("file:///a.md")]])
|
||||||
|
poller = _periodic(source, cfg, jobs, sync, default_max_attempts=3)
|
||||||
|
await poller._sweep_once()
|
||||||
|
queued = await jobs.list_jobs(source_id="src")
|
||||||
|
assert queued[0].max_attempts == 9
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_storage_options_thread_through_to_job_extra(jobs, sync):
|
||||||
|
cfg = S3SourceConfig(
|
||||||
|
type="s3",
|
||||||
|
id="bucket",
|
||||||
|
uri="s3://bucket/",
|
||||||
|
storage_options={"endpoint": "http://seaweed:8333"},
|
||||||
|
)
|
||||||
|
source = _StubSource(
|
||||||
|
"bucket", [[_event("s3://bucket/file.md", source_id="bucket")]]
|
||||||
|
)
|
||||||
|
poller = _periodic(source, cfg, jobs, sync)
|
||||||
|
await poller._sweep_once()
|
||||||
|
queued = await jobs.list_jobs(source_id="bucket")
|
||||||
|
assert queued[0].extra == {"storage_options": {"endpoint": "http://seaweed:8333"}}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_http_headers_thread_through_to_job_extra(jobs, sync):
|
||||||
|
cfg = HTTPSourceConfig(
|
||||||
|
type="http",
|
||||||
|
id="auth",
|
||||||
|
urls=["https://example.com/a.md"],
|
||||||
|
headers={"Authorization": "Bearer abc"},
|
||||||
|
)
|
||||||
|
source = _StubSource(
|
||||||
|
"auth", [[_event("https://example.com/a.md", source_id="auth")]]
|
||||||
|
)
|
||||||
|
poller = _periodic(source, cfg, jobs, sync)
|
||||||
|
await poller._sweep_once()
|
||||||
|
queued = await jobs.list_jobs(source_id="auth")
|
||||||
|
assert queued[0].extra == {"headers": {"Authorization": "Bearer abc"}}
|
||||||
|
|
||||||
|
|
||||||
|
# --- PollerManager lifecycle ---
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_manager_builds_pollers_per_source(tmp_path, jobs, sync):
|
||||||
|
"""When SourceConfig.id is set, the poller's source uses it verbatim;
|
||||||
|
when omitted, the adapter auto-derives one from its target."""
|
||||||
|
configs = [
|
||||||
|
FSSourceConfig(type="fs", root=tmp_path),
|
||||||
|
S3SourceConfig(type="s3", uri="s3://bucket/"),
|
||||||
|
HTTPSourceConfig(type="http", id="urls", urls=[]),
|
||||||
|
]
|
||||||
|
manager = PollerManager(
|
||||||
|
configs=configs,
|
||||||
|
job_repo=jobs,
|
||||||
|
sync_repo=sync,
|
||||||
|
)
|
||||||
|
built = manager.build_pollers()
|
||||||
|
assert len(built) == 3
|
||||||
|
assert {p.source_id for p in built} == {
|
||||||
|
f"fs:{tmp_path.resolve()}",
|
||||||
|
"s3:bucket/",
|
||||||
|
"urls",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_manager_double_start_raises(tmp_path, jobs, sync):
|
||||||
|
cfg = FSSourceConfig(
|
||||||
|
type="fs",
|
||||||
|
id="local",
|
||||||
|
root=tmp_path,
|
||||||
|
poll_interval_s=60.0,
|
||||||
|
)
|
||||||
|
manager = PollerManager(configs=[cfg], job_repo=jobs, sync_repo=sync)
|
||||||
|
await manager.start()
|
||||||
|
try:
|
||||||
|
with pytest.raises(RuntimeError, match="already started"):
|
||||||
|
await manager.start()
|
||||||
|
finally:
|
||||||
|
await manager.stop()
|
||||||
|
|
||||||
|
|
||||||
|
# --- FSPoller end-to-end smoke ---
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_fs_poller_enqueues_initial_files(tmp_path, jobs, sync):
|
||||||
|
(tmp_path / "a.md").write_text("hello")
|
||||||
|
(tmp_path / "b.md").write_text("world")
|
||||||
|
|
||||||
|
cfg = FSSourceConfig(type="fs", id="local", root=tmp_path, poll_interval_s=60.0)
|
||||||
|
manager = PollerManager(
|
||||||
|
configs=[cfg], job_repo=jobs, sync_repo=sync, supported_extensions=[".md"]
|
||||||
|
)
|
||||||
|
await manager.start()
|
||||||
|
try:
|
||||||
|
# Wait for the initial sweep to land jobs.
|
||||||
|
for _ in range(40):
|
||||||
|
queued = await jobs.list_jobs(source_id=f"fs:{tmp_path.resolve()}")
|
||||||
|
if len(queued) == 2:
|
||||||
|
break
|
||||||
|
await asyncio.sleep(0.05)
|
||||||
|
finally:
|
||||||
|
await manager.stop()
|
||||||
|
|
||||||
|
queued = await jobs.list_jobs(source_id=f"fs:{tmp_path.resolve()}")
|
||||||
|
assert {Path(j.uri).name for j in queued} == {"a.md", "b.md"}
|
||||||
|
assert all(j.status is JobStatus.QUEUED for j in queued)
|
||||||
Loading…
Reference in a new issue